use crate::{Request, Response, Result, DnsError};
use crate::types::{EdnsRecord, EdnsOption, edns_option_codes};
use super::{Transport, TransportConfig};
use async_trait::async_trait;
use std::time::Duration;
use tokio::net::UdpSocket;
use tokio::time::timeout;
use crate::{dns_debug, dns_info, dns_error, dns_transport};
#[derive(Debug)]
pub struct UdpTransport {
config: TransportConfig,
}
impl UdpTransport {
pub fn new(config: TransportConfig) -> Self {
Self { config }
}
#[cfg(windows)]
async fn create_windows_socket(&self) -> Result<UdpSocket> {
use std::net::SocketAddr;
dns_debug!("Windows平台:开始创建UDP socket");
let bind_addresses = [
"0.0.0.0:0",
"127.0.0.1:0",
"::1:0", ];
let mut last_error = None;
for addr in &bind_addresses {
dns_debug!("尝试绑定地址: {}", addr);
match UdpSocket::bind(addr).await {
Ok(socket) => {
dns_debug!("成功绑定到地址: {}", addr);
return Ok(socket);
}
Err(e) => {
dns_debug!("绑定失败 {}: {}", addr, e);
last_error = Some(e);
continue;
}
}
}
dns_error!("所有绑定尝试都失败");
if let Some(e) = last_error {
Err(DnsError::Network(format!("Windows UDP socket 绑定失败: {}", e)))
} else {
Err(DnsError::Network("Windows UDP socket 绑定失败: 未知错误".to_string()))
}
}
#[cfg(windows)]
async fn configure_windows_socket(&self, socket: &UdpSocket) -> Result<()> {
Ok(())
}
#[cfg(not(windows))]
async fn create_windows_socket(&self) -> Result<UdpSocket> {
Err(DnsError::Network("在非Windows平台调用create_windows_socket方法".to_string()))
}
#[cfg(not(windows))]
async fn configure_windows_socket(&self, _socket: &UdpSocket) -> Result<()> {
Ok(())
}
pub fn serialize_request(request: &Request) -> Result<Vec<u8>> {
dns_debug!("开始序列化DNS请求");
dns_debug!("请求ID: {}", request.id);
dns_debug!("查询域名: '{}'", request.query.name);
dns_debug!("查询类型: {:?}", request.query.qtype);
dns_debug!("客户端地址: {:?}", request.client_address);
let mut buffer = Vec::with_capacity(512);
let has_edns = request.client_address.is_some();
let additional_count = if has_edns { 1u16 } else { 0u16 };
dns_debug!("需要EDNS记录: {}, 附加记录数: {}", has_edns, additional_count);
buffer.extend_from_slice(&request.id.to_be_bytes());
let mut flags = 0u16;
if request.flags.qr { flags |= 0x8000; }
flags |= (request.flags.opcode as u16) << 11;
if request.flags.aa { flags |= 0x0400; }
if request.flags.tc { flags |= 0x0200; }
if request.flags.rd { flags |= 0x0100; }
if request.flags.ra { flags |= 0x0080; }
flags |= (request.flags.z as u16) << 4;
flags |= request.flags.rcode as u16;
buffer.extend_from_slice(&flags.to_be_bytes());
dns_debug!("DNS头部标志位: 0x{:04X}", flags);
buffer.extend_from_slice(&1u16.to_be_bytes());
buffer.extend_from_slice(&0u16.to_be_bytes());
buffer.extend_from_slice(&0u16.to_be_bytes());
buffer.extend_from_slice(&additional_count.to_be_bytes());
dns_debug!("DNS头部完成,当前缓冲区长度: {} 字节", buffer.len());
let name_start_pos = buffer.len();
Self::encode_name(&request.query.name, &mut buffer)?;
let name_end_pos = buffer.len();
dns_debug!("域名编码完成,占用 {} 字节 (位置 {}-{})", name_end_pos - name_start_pos, name_start_pos, name_end_pos);
buffer.extend_from_slice(&u16::from(request.query.qtype).to_be_bytes());
buffer.extend_from_slice(&u16::from(request.query.qclass).to_be_bytes());
dns_debug!("查询类型和类别添加完成,当前缓冲区长度: {} 字节", buffer.len());
if let Some(ref client_address) = request.client_address {
dns_debug!("添加EDNS记录");
Self::encode_edns_record(&mut buffer, client_address)?;
dns_debug!("EDNS记录添加完成,最终缓冲区长度: {} 字节", buffer.len());
}
dns_debug!("DNS请求序列化完成,总长度: {} 字节", buffer.len());
let preview_len = buffer.len().min(64);
let hex_preview: String = buffer[..preview_len].iter()
.map(|b| format!("{:02X}", b))
.collect::<Vec<_>>()
.join(" ");
dns_debug!("请求数据预览 (前{}字节): {}", preview_len, hex_preview);
Ok(buffer)
}
pub fn encode_name(name: &str, buffer: &mut Vec<u8>) -> Result<()> {
dns_debug!("编码域名: '{}'", name);
if name.is_empty() || name == "." {
dns_debug!("空域名或根域名,添加终止符");
buffer.push(0);
return Ok(());
}
let name = name.trim_end_matches('.');
dns_debug!("处理后的域名: '{}'", name);
for (i, label) in name.split('.').enumerate() {
if label.is_empty() {
dns_debug!("跳过空标签 {}", i);
continue;
}
if label.len() > 63 {
dns_debug!("标签 '{}' 长度超过63字节", label);
return Err(DnsError::Protocol("标签长度过长".to_string()));
}
dns_debug!("添加标签 {}: '{}' (长度: {})", i, label, label.len());
buffer.push(label.len() as u8);
buffer.extend_from_slice(label.as_bytes());
}
dns_debug!("添加域名终止符");
buffer.push(0);
dns_debug!("域名编码完成,总长度: {} 字节", buffer.len());
Ok(())
}
pub fn serialize_response(response: &Response) -> Result<Vec<u8>> {
dns_debug!("开始序列化DNS响应");
dns_debug!("响应ID: {}", response.id);
dns_debug!("查询数: {}, 回答数: {}, 权威数: {}, 附加数: {}",
response.queries.len(), response.answers.len(),
response.authorities.len(), response.additionals.len());
let mut buffer = Vec::with_capacity(512);
buffer.extend_from_slice(&response.id.to_be_bytes());
let mut flags = 0u16;
if response.flags.qr { flags |= 0x8000; }
flags |= (response.flags.opcode as u16) << 11;
if response.flags.aa { flags |= 0x0400; }
if response.flags.tc { flags |= 0x0200; }
if response.flags.rd { flags |= 0x0100; }
if response.flags.ra { flags |= 0x0080; }
flags |= (response.flags.z as u16) << 4;
flags |= response.flags.rcode as u16;
buffer.extend_from_slice(&flags.to_be_bytes());
dns_debug!("DNS头部标志位: 0x{:04X}", flags);
buffer.extend_from_slice(&(response.queries.len() as u16).to_be_bytes());
buffer.extend_from_slice(&(response.answers.len() as u16).to_be_bytes());
buffer.extend_from_slice(&(response.authorities.len() as u16).to_be_bytes());
buffer.extend_from_slice(&(response.additionals.len() as u16).to_be_bytes());
dns_debug!("DNS头部完成,当前缓冲区长度: {} 字节", buffer.len());
for query in &response.queries {
Self::encode_name(&query.name, &mut buffer)?;
buffer.extend_from_slice(&u16::from(query.qtype).to_be_bytes());
buffer.extend_from_slice(&u16::from(query.qclass).to_be_bytes());
}
dns_debug!("查询部分序列化完成,当前缓冲区长度: {} 字节", buffer.len());
for record in &response.answers {
Self::encode_record(record, &mut buffer)?;
}
dns_debug!("回答部分序列化完成,当前缓冲区长度: {} 字节", buffer.len());
for record in &response.authorities {
Self::encode_record(record, &mut buffer)?;
}
dns_debug!("权威部分序列化完成,当前缓冲区长度: {} 字节", buffer.len());
for record in &response.additionals {
Self::encode_record(record, &mut buffer)?;
}
dns_debug!("附加部分序列化完成,最终缓冲区长度: {} 字节", buffer.len());
let preview_len = buffer.len().min(64);
let hex_preview: String = buffer[..preview_len].iter()
.map(|b| format!("{:02X}", b))
.collect::<Vec<_>>()
.join(" ");
dns_debug!("响应数据预览 (前{}字节): {}", preview_len, hex_preview);
Ok(buffer)
}
pub fn encode_record(record: &crate::types::Record, buffer: &mut Vec<u8>) -> Result<()> {
Self::encode_name(&record.name, buffer)?;
buffer.extend_from_slice(&u16::from(record.rtype).to_be_bytes());
buffer.extend_from_slice(&u16::from(record.class).to_be_bytes());
buffer.extend_from_slice(&record.ttl.to_be_bytes());
let data_bytes = Self::encode_record_data(&record.data)?;
buffer.extend_from_slice(&(data_bytes.len() as u16).to_be_bytes());
buffer.extend_from_slice(&data_bytes);
Ok(())
}
pub fn encode_record_data(data: &crate::types::RecordData) -> Result<Vec<u8>> {
use crate::types::RecordData;
match data {
RecordData::A(ip) => Ok(ip.octets().to_vec()),
RecordData::AAAA(ip) => Ok(ip.octets().to_vec()),
RecordData::CNAME(name) | RecordData::NS(name) | RecordData::PTR(name) => {
let mut buffer = Vec::new();
Self::encode_name(name, &mut buffer)?;
Ok(buffer)
},
RecordData::MX { priority, exchange } => {
let mut buffer = Vec::new();
buffer.extend_from_slice(&priority.to_be_bytes());
Self::encode_name(exchange, &mut buffer)?;
Ok(buffer)
},
RecordData::TXT(texts) => {
let mut buffer = Vec::new();
for text in texts {
if text.len() > 255 {
return Err(DnsError::Protocol("TXT记录长度过长".to_string()));
}
buffer.push(text.len() as u8);
buffer.extend_from_slice(text.as_bytes());
}
Ok(buffer)
},
RecordData::SOA { mname, rname, serial, refresh, retry, expire, minimum } => {
let mut buffer = Vec::new();
Self::encode_name(mname, &mut buffer)?;
Self::encode_name(rname, &mut buffer)?;
buffer.extend_from_slice(&serial.to_be_bytes());
buffer.extend_from_slice(&refresh.to_be_bytes());
buffer.extend_from_slice(&retry.to_be_bytes());
buffer.extend_from_slice(&expire.to_be_bytes());
buffer.extend_from_slice(&minimum.to_be_bytes());
Ok(buffer)
},
RecordData::SRV { priority, weight, port, target } => {
let mut buffer = Vec::new();
buffer.extend_from_slice(&priority.to_be_bytes());
buffer.extend_from_slice(&weight.to_be_bytes());
buffer.extend_from_slice(&port.to_be_bytes());
Self::encode_name(target, &mut buffer)?;
Ok(buffer)
},
RecordData::Unknown(data) => Ok(data.clone()),
}
}
pub fn deserialize_request(data: &[u8]) -> Result<Request> {
if data.len() < 12 {
return Err(DnsError::Protocol("请求数据过短".to_string()));
}
let id = u16::from_be_bytes([data[0], data[1]]);
let flags_raw = u16::from_be_bytes([data[2], data[3]]);
let flags = crate::types::Flags {
qr: (flags_raw & 0x8000) != 0,
opcode: ((flags_raw >> 11) & 0x0F) as u8,
aa: (flags_raw & 0x0400) != 0,
tc: (flags_raw & 0x0200) != 0,
rd: (flags_raw & 0x0100) != 0,
ra: (flags_raw & 0x0080) != 0,
z: ((flags_raw >> 4) & 0x07) as u8,
rcode: (flags_raw & 0x0F) as u8,
};
let qdcount = u16::from_be_bytes([data[4], data[5]]);
if qdcount != 1 {
return Err(DnsError::Protocol("请求必须包含且仅包含一个查询".to_string()));
}
let mut offset = 12;
let (query, _) = Self::parse_query(data, offset)?;
let client_address = None;
Ok(Request {
id,
flags,
query,
client_address,
})
}
pub fn deserialize_response(data: &[u8]) -> Result<Response> {
if data.len() < 12 {
return Err(DnsError::Protocol("响应数据过短".to_string()));
}
let id = u16::from_be_bytes([data[0], data[1]]);
let flags_raw = u16::from_be_bytes([data[2], data[3]]);
let flags = crate::types::Flags {
qr: (flags_raw & 0x8000) != 0,
opcode: ((flags_raw >> 11) & 0x0F) as u8,
aa: (flags_raw & 0x0400) != 0,
tc: (flags_raw & 0x0200) != 0,
rd: (flags_raw & 0x0100) != 0,
ra: (flags_raw & 0x0080) != 0,
z: ((flags_raw >> 4) & 0x07) as u8,
rcode: (flags_raw & 0x0F) as u8,
};
let qdcount = u16::from_be_bytes([data[4], data[5]]);
let ancount = u16::from_be_bytes([data[6], data[7]]);
let nscount = u16::from_be_bytes([data[8], data[9]]);
let arcount = u16::from_be_bytes([data[10], data[11]]);
let mut offset = 12;
let mut queries = Vec::new();
let mut answers = Vec::new();
let mut authorities = Vec::new();
let mut additionals = Vec::new();
for _ in 0..qdcount {
let (query, new_offset) = Self::parse_query(data, offset)?;
queries.push(query);
offset = new_offset;
}
for _ in 0..ancount {
let (record, new_offset) = Self::parse_record(data, offset)?;
answers.push(record);
offset = new_offset;
}
for _ in 0..nscount {
let (record, new_offset) = Self::parse_record(data, offset)?;
authorities.push(record);
offset = new_offset;
}
for _ in 0..arcount {
let (record, new_offset) = Self::parse_record(data, offset)?;
additionals.push(record);
offset = new_offset;
}
Ok(Response {
id,
flags,
queries,
answers,
authorities,
additionals,
})
}
pub fn parse_query(data: &[u8], offset: usize) -> Result<(crate::types::Query, usize)> {
let (name, mut offset) = Self::parse_name(data, offset)?;
if offset + 4 > data.len() {
return Err(DnsError::Protocol("查询格式无效".to_string()));
}
let qtype = u16::from_be_bytes([data[offset], data[offset + 1]]).into();
let qclass = u16::from_be_bytes([data[offset + 2], data[offset + 3]]).into();
offset += 4;
Ok((crate::types::Query { name, qtype, qclass }, offset))
}
pub fn parse_record(data: &[u8], offset: usize) -> Result<(crate::types::Record, usize)> {
let (name, mut offset) = Self::parse_name(data, offset)?;
if offset + 10 > data.len() {
return Err(DnsError::Protocol("记录格式无效".to_string()));
}
let rtype = u16::from_be_bytes([data[offset], data[offset + 1]]).into();
let class = u16::from_be_bytes([data[offset + 2], data[offset + 3]]).into();
let ttl = u32::from_be_bytes([data[offset + 4], data[offset + 5], data[offset + 6], data[offset + 7]]);
let rdlength = u16::from_be_bytes([data[offset + 8], data[offset + 9]]) as usize;
offset += 10;
if offset + rdlength > data.len() {
return Err(DnsError::Protocol("记录数据长度无效".to_string()))
}
let rdata = &data[offset..offset + rdlength];
let record_data = Self::parse_record_data(rtype, rdata, data, offset)?;
offset += rdlength;
Ok((crate::types::Record {
name,
rtype,
class,
ttl,
data: record_data,
}, offset))
}
pub fn parse_name(data: &[u8], mut offset: usize) -> Result<(String, usize)> {
dns_debug!("开始解析域名,起始偏移: {}, 数据长度: {}", offset, data.len());
let mut name = String::new();
let mut jumped = false;
let mut jump_offset = 0;
let mut loop_count = 0;
const MAX_LOOPS: usize = 100;
loop {
loop_count += 1;
if loop_count > MAX_LOOPS {
dns_debug!("域名解析循环次数超限,可能存在循环引用");
return Err(DnsError::Protocol("域名解析检测到循环引用".to_string()));
}
if offset >= data.len() {
dns_debug!("偏移量 {} 超出数据长度 {}", offset, data.len());
return Err(DnsError::Protocol("域名解析数据溢出".to_string()));
}
let len = data[offset];
dns_debug!("偏移 {}: 长度字节 = 0x{:02X} ({})", offset, len, len);
if len == 0 {
dns_debug!("遇到域名终止符,解析完成");
offset += 1;
break;
}
if (len & 0xC0) == 0xC0 {
if offset + 1 >= data.len() {
dns_debug!("压缩指针数据不完整");
return Err(DnsError::Protocol("压缩指针数据不完整".to_string()));
}
let pointer = (((len & 0x3F) as usize) << 8) | (data[offset + 1] as usize);
dns_debug!("压缩指针指向偏移: {}", pointer);
if pointer >= data.len() {
dns_debug!("压缩指针 {} 超出数据范围 {}", pointer, data.len());
return Err(DnsError::Protocol("压缩指针无效".to_string()));
}
if !jumped {
jump_offset = offset + 2;
jumped = true;
dns_debug!("设置跳转返回点: {}", jump_offset);
}
offset = pointer;
continue;
}
if len > 63 {
dns_debug!("标签长度 {} 超过63字节限制", len);
return Err(DnsError::Protocol("标签长度过长".to_string()));
}
offset += 1;
if offset + len as usize > data.len() {
dns_debug!("标签数据超出范围: 偏移{}+长度{} > 数据长度{}", offset, len, data.len());
return Err(DnsError::Protocol("域名标签数据溢出".to_string()));
}
if !name.is_empty() {
name.push('.');
}
let label = String::from_utf8_lossy(&data[offset..offset + len as usize]);
dns_debug!("解析标签: '{}'", label);
name.push_str(&label);
offset += len as usize;
}
if jumped {
offset = jump_offset;
dns_debug!("恢复到跳转返回点: {}", offset);
}
dns_debug!("域名解析完成: '{}', 最终偏移: {}", name, offset);
Ok((name, offset))
}
pub fn parse_record_data(
rtype: crate::types::RecordType,
rdata: &[u8],
full_data: &[u8],
rdata_offset: usize,
) -> Result<crate::types::RecordData> {
use crate::types::{RecordType, RecordData};
use std::net::{Ipv4Addr, Ipv6Addr};
match rtype {
RecordType::A => {
if rdata.len() != 4 {
return Err(DnsError::Protocol("A记录长度无效".to_string()));
}
Ok(RecordData::A(Ipv4Addr::new(rdata[0], rdata[1], rdata[2], rdata[3])))
}
RecordType::AAAA => {
if rdata.len() != 16 {
return Err(DnsError::Protocol("AAAA记录长度无效".to_string()));
}
let mut addr = [0u8; 16];
addr.copy_from_slice(rdata);
Ok(RecordData::AAAA(Ipv6Addr::from(addr)))
}
RecordType::CNAME | RecordType::NS | RecordType::PTR => {
let (name, _) = Self::parse_name(full_data, rdata_offset)?;
match rtype {
RecordType::CNAME => Ok(RecordData::CNAME(name)),
RecordType::NS => Ok(RecordData::NS(name)),
RecordType::PTR => Ok(RecordData::PTR(name)),
_ => unreachable!(),
}
}
RecordType::MX => {
if rdata.len() < 3 {
return Err(DnsError::Protocol("MX记录长度无效".to_string()));
}
let priority = u16::from_be_bytes([rdata[0], rdata[1]]);
let (exchange, _) = Self::parse_name(full_data, rdata_offset + 2)?;
Ok(RecordData::MX { priority, exchange })
}
RecordType::TXT => {
let mut texts = Vec::new();
let mut offset = 0;
while offset < rdata.len() {
if offset >= rdata.len() {
break;
}
let len = rdata[offset] as usize;
offset += 1;
if offset + len > rdata.len() {
return Err(DnsError::Protocol("TXT记录格式无效".to_string()))
}
let text = String::from_utf8_lossy(&rdata[offset..offset + len]).to_string();
texts.push(text);
offset += len;
}
Ok(RecordData::TXT(texts))
}
_ => Ok(RecordData::Unknown(rdata.to_vec())),
}
}
pub fn encode_edns_record(buffer: &mut Vec<u8>, client_address: &crate::types::ClientAddress) -> Result<()> {
buffer.push(0x00);
buffer.extend_from_slice(&41u16.to_be_bytes());
buffer.extend_from_slice(&4096u16.to_be_bytes());
buffer.push(0); buffer.push(0); buffer.extend_from_slice(&0u16.to_be_bytes());
let client_address_data = client_address.encode();
let option_length = client_address_data.len() as u16;
let rdlength = 4 + option_length;
buffer.extend_from_slice(&rdlength.to_be_bytes());
buffer.extend_from_slice(&edns_option_codes::CLIENT_ADDRESS.to_be_bytes());
buffer.extend_from_slice(&option_length.to_be_bytes());
buffer.extend_from_slice(&client_address_data);
Ok(())
}
}
#[async_trait]
impl Transport for UdpTransport {
async fn send(&self, request: &Request) -> Result<Response> {
dns_debug!("UDP传输开始发送请求");
dns_debug!("目标域名: {}", request.query.name);
dns_debug!("查询类型: {:?}", request.query.qtype);
let socket = if cfg!(windows) {
dns_debug!("使用Windows平台socket创建策略");
self.create_windows_socket().await?
} else {
dns_debug!("使用Unix/Linux平台socket创建策略");
UdpSocket::bind("0.0.0.0:0").await
.map_err(|e| DnsError::Network(format!("UDP socket 绑定失败: {}", e)))?
};
let server_addr = format!("{}:{}", self.config.server, self.config.port);
dns_debug!("DNS服务器地址: {}", server_addr);
if cfg!(windows) {
dns_debug!("配置Windows socket选项");
self.configure_windows_socket(&socket).await?;
}
let request_data = Self::serialize_request(request)?;
dns_debug!("请求数据长度: {} 字节", request_data.len());
let send_result = timeout(
self.config.timeout,
socket.send_to(&request_data, &server_addr)
).await;
match send_result {
Ok(Ok(_)) => {},
Ok(Err(e)) => {
let error_msg = if cfg!(windows) {
format!("Windows UDP 发送失败: {} (服务器: {})", e, server_addr)
} else {
format!("UDP 发送失败: {} (服务器: {})", e, server_addr)
};
return Err(DnsError::Network(error_msg));
},
Err(_) => return Err(DnsError::Timeout),
}
let mut buffer = [0u8; 512];
let recv_result = timeout(
self.config.timeout,
socket.recv(&mut buffer)
).await;
let len = match recv_result {
Ok(Ok(len)) => len,
Ok(Err(e)) => {
let error_msg = if cfg!(windows) {
format!("Windows UDP 接收失败: {}", e)
} else {
format!("UDP 接收失败: {}", e)
};
return Err(DnsError::Network(error_msg));
},
Err(_) => return Err(DnsError::Timeout),
};
dns_debug!("收到DNS响应,长度: {} 字节", len);
let preview_len = len.min(64);
let hex_preview: String = buffer[..preview_len].iter()
.map(|b| format!("{:02X}", b))
.collect::<Vec<_>>()
.join(" ");
dns_debug!("响应数据预览 (前{}字节): {}", preview_len, hex_preview);
let result = Self::deserialize_response(&buffer[..len]);
match &result {
Ok(response) => {
dns_debug!("DNS响应解析成功,包含 {} 个回答记录", response.answers.len());
},
Err(e) => {
dns_error!("DNS响应解析失败: {}", e);
}
}
result
}
fn transport_type(&self) -> &'static str {
"UDP"
}
fn set_timeout(&mut self, timeout: Duration) {
self.config.timeout = timeout;
}
fn timeout(&self) -> Duration {
self.config.timeout
}
}