Files
vnt/vnt-core/src/utils/dns_query.rs
T

314 lines
10 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
use anyhow::{Context, anyhow};
use dns_parser::{Builder, Packet, QueryClass, QueryType, RData, ResponseCode};
use rand::seq::SliceRandom;
use std::io;
use std::net::{Ipv4Addr, Ipv6Addr, SocketAddr};
use std::str::FromStr;
use std::time::Duration;
use tokio::net::UdpSocket;
use rust_p2p_core::socket::LocalInterface;
pub async fn dns_query_txt(
domain: &str,
mut name_servers: Vec<String>,
default_interface: &Option<LocalInterface>,
) -> io::Result<Vec<String>> {
let mut err: Option<io::Error> = None;
if name_servers.is_empty() {
name_servers.push("223.5.5.5:53".into());
name_servers.push("114.114.114.114:53".into());
}
for name_server in name_servers {
match txt_dns(domain, name_server, default_interface).await {
Ok(addr) => {
if !addr.is_empty() {
return Ok(addr);
}
}
Err(e) => {
err.replace(e);
}
}
continue;
}
if let Some(e) = err {
Err(e)
} else {
Err(io::Error::other(format!("DNS query failed {domain:?}")))
}
}
pub async fn dns_query_one(
domain: &str,
name_servers: &Vec<String>,
default_interface: &Option<LocalInterface>,
) -> anyhow::Result<SocketAddr> {
let mut vec = dns_query_all(domain, name_servers, default_interface).await?;
if default_interface.is_some() {
// 出口网卡绑定目前应用于 IPv4 Socket;优先且只使用可绑定的 IPv4 地址。
vec.retain(SocketAddr::is_ipv4);
}
vec.shuffle(&mut rand::rng());
vec.pop().context("DNS query failed")
}
pub async fn dns_query_all(
domain: &str,
name_servers: &Vec<String>,
default_interface: &Option<LocalInterface>,
) -> anyhow::Result<Vec<SocketAddr>> {
match SocketAddr::from_str(domain) {
Ok(addr) => Ok(vec![addr]),
Err(_) => {
if name_servers.is_empty() {
let addrs: Vec<SocketAddr> = tokio::net::lookup_host(domain)
.await
.map_err(|e| io::Error::other(format!("DNS query failed: {domain:?},{e:?}")))?
.collect();
return Ok(addrs);
}
let mut err: Option<io::Error> = None;
for name_server in name_servers {
let end_index = domain
.rfind(':')
.ok_or_else(|| io::Error::other(format!("not port: {domain:?}")))?;
let host = &domain[..end_index];
let port = u16::from_str(&domain[end_index + 1..])
.map_err(|_| io::Error::other(format!("not port: {domain:?}")))?;
let th1 = {
let host = host.to_string();
let name_server = name_server.clone();
let default_interface = default_interface.clone();
tokio::spawn(a_dns(host, name_server, default_interface.clone()))
};
let th2 = {
let host = host.to_string();
let name_server = name_server.clone();
let default_interface = default_interface.clone();
tokio::spawn(aaaa_dns(host, name_server, default_interface.clone()))
};
let mut addr = Vec::new();
match th1.await? {
Ok(rs) => {
for ip in rs {
addr.push(SocketAddr::new(ip.into(), port));
}
}
Err(e) => {
err.replace(e);
}
}
match th2.await? {
Ok(rs) => {
for ip in rs {
addr.push(SocketAddr::new(ip.into(), port));
}
}
Err(e) => {
if addr.is_empty() {
err.replace(e);
continue;
}
}
}
if addr.is_empty() {
continue;
}
return Ok(addr);
}
if let Some(e) = err {
Err(e.into())
} else {
Err(anyhow!("DNS query failed {domain:?}"))
}
}
}
}
/// 校验域名格式:label 非空且不超过 63 字节,全长不超过 253 字节。
/// dns-parser 的 add_question 对非法 label 直接 assert panic
/// 必须在调用前拦截
fn is_valid_domain(domain: &str) -> bool {
let domain = domain.strip_suffix('.').unwrap_or(domain);
!domain.is_empty()
&& domain.len() <= 253
&& domain
.split('.')
.all(|label| !label.is_empty() && label.len() <= 63)
}
async fn query<'a>(
udp: &UdpSocket,
domain: &str,
name_server: SocketAddr,
record_type: QueryType,
buf: &'a mut [u8],
) -> io::Result<Packet<'a>> {
if !is_valid_domain(domain) {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
format!("invalid domain {domain:?}"),
));
}
let mut builder = Builder::new_query(1, true);
builder.add_question(domain, false, record_type, QueryClass::IN);
// 非法域名(如 label 超长)build 会失败,不能 unwrap panic
let packet = builder.build().map_err(|e| {
io::Error::new(
io::ErrorKind::InvalidInput,
format!("invalid domain {domain:?}: {e:?}"),
)
})?;
udp.connect(name_server).await?;
let mut count = 0;
let len = loop {
udp.send(&packet).await?;
match tokio::time::timeout(Duration::from_secs(3), udp.recv(buf)).await {
Ok(len) => {
break len?;
}
Err(_) => {
count += 1;
if count < 3 {
continue;
}
Err(io::Error::other(format!("DNS {name_server:?} recv error ")))?
}
};
};
let pkt = Packet::parse(&buf[..len]).map_err(|e| {
io::Error::other(format!(
"domain {domain:?} DNS {name_server:?} data error: {e}"
))
})?;
if pkt.header.response_code != ResponseCode::NoError {
return Err(io::Error::other(format!(
"response_code {} DNS {:?} domain {:?}",
pkt.header.response_code, name_server, domain
)));
}
if pkt.answers.is_empty() {
return Err(io::Error::other(format!(
"No records received DNS {name_server:?} domain {domain:?}"
)));
}
Ok(pkt)
}
pub async fn txt_dns(
domain: &str,
name_server: String,
default_interface: &Option<LocalInterface>,
) -> io::Result<Vec<String>> {
let name_server: SocketAddr = name_server
.parse()
.map_err(|e| io::Error::other(format!("dns {name_server} is error :{e:?}")))?;
let udp = bind_udp(name_server, default_interface)?;
let mut buf = vec![0u8; 65536];
let message = query(&udp, domain, name_server, QueryType::TXT, &mut buf).await?;
let mut rs = Vec::new();
for record in message.answers {
if let RData::TXT(txt) = record.data {
for x in txt.iter() {
let txt = std::str::from_utf8(x)
.map_err(|_| io::Error::other("record type txt is not string"))?;
rs.push(txt.to_string());
}
}
}
Ok(rs)
}
fn bind_udp(
name_server: SocketAddr,
default_interface: &Option<LocalInterface>,
) -> io::Result<UdpSocket> {
let addr: SocketAddr = if name_server.is_ipv4() {
SocketAddr::from(([0, 0, 0, 0], 0))
} else {
SocketAddr::from(([0; 8], 0))
};
let socket = rust_p2p_core::socket::bind_udp(addr, default_interface.as_ref())?;
UdpSocket::from_std(socket.into())
}
pub async fn a_dns(
domain: String,
name_server: String,
default_interface: Option<LocalInterface>,
) -> io::Result<Vec<Ipv4Addr>> {
let name_server: SocketAddr = name_server
.parse()
.map_err(|e| io::Error::other(format!("dns {name_server} is error :{e:?}")))?;
let udp = bind_udp(name_server, &default_interface)?;
let mut buf = vec![0u8; 65536];
let message = query(&udp, &domain, name_server, QueryType::A, &mut buf).await?;
let mut rs = Vec::new();
for record in message.answers {
if let RData::A(a) = record.data {
rs.push(a.0);
}
}
Ok(rs)
}
pub async fn aaaa_dns(
domain: String,
name_server: String,
default_interface: Option<LocalInterface>,
) -> io::Result<Vec<Ipv6Addr>> {
let name_server: SocketAddr = name_server
.parse()
.map_err(|e| io::Error::other(format!("dns {name_server} is error :{e:?}")))?;
let udp = bind_udp(name_server, &default_interface)?;
let mut buf = vec![0u8; 65536];
let message = query(&udp, &domain, name_server, QueryType::AAAA, &mut buf).await?;
let mut rs = Vec::new();
for record in message.answers {
if let RData::AAAA(a) = record.data {
rs.push(a.0);
}
}
Ok(rs)
}
#[cfg(test)]
mod tests {
use super::*;
/// 非法域名(label 超过 63 字节)必须返回错误而不是 panic
#[tokio::test]
async fn test_query_invalid_domain_no_panic() {
let udp = UdpSocket::bind("0.0.0.0:0").await.unwrap();
let mut buf = vec![0u8; 512];
let bad_domain = format!("{}.com", "a".repeat(64));
let rs = query(
&udp,
&bad_domain,
"127.0.0.1:53".parse().unwrap(),
QueryType::A,
&mut buf,
)
.await;
let err = rs.expect_err("invalid domain must be rejected");
assert_eq!(err.kind(), io::ErrorKind::InvalidInput);
}
#[test]
fn test_is_valid_domain() {
assert!(is_valid_domain("example.com"));
assert!(is_valid_domain("a-b_1.example.com"));
assert!(is_valid_domain("example.com.")); // FQDN 尾点合法
assert!(is_valid_domain(&format!("{}.com", "a".repeat(63))));
assert!(!is_valid_domain(""));
assert!(!is_valid_domain(&format!("{}.com", "a".repeat(64))));
assert!(!is_valid_domain("a..b"));
assert!(!is_valid_domain(&"a".repeat(254)));
}
}