use std::sync::Arc;
use std::time::Instant;
use uuid::Uuid;
use crate::resolver::{CoreResolverConfig, CoreResolver};
use crate::upstream_handler::UpstreamManager;
use crate::utils::{parse_simple_server_address, parse_url_components, get_user_agent};
use crate::error::{DnsError, Result};
use crate::{dns_info, dns_debug};
use super::{
strategy::QueryStrategy,
engine::SmartDecisionEngine,
types::{DnsQueryRequest, DnsQueryResponse, DnsRecord, DnsRecordType},
};
#[derive(Debug)]
pub struct SmartDnsResolver {
resolver: CoreResolver,
upstream_manager: UpstreamManager,
decision_engine: Option<Arc<SmartDecisionEngine>>,
query_strategy: QueryStrategy,
enable_edns: bool,
}
impl Drop for SmartDnsResolver {
fn drop(&mut self) {
dns_info!("Dropping SmartDnsResolver, cleaning up resources...");
dns_debug!("SmartDnsResolver dropped with {} transports", self.resolver.transport_count());
}
}
impl SmartDnsResolver {
pub(super) fn new(
config: CoreResolverConfig,
upstream_manager: UpstreamManager,
decision_engine: Option<Arc<SmartDecisionEngine>>,
query_strategy: QueryStrategy,
enable_edns: bool,
) -> Result<Self> {
let default_timeout = config.default_timeout;
let mut resolver = CoreResolver::new(config);
let specs = upstream_manager.get_specs();
dns_debug!("SmartDnsResolver::new - 开始处理 {} 个上游服务器", specs.len());
for spec in specs {
match spec.transport_type {
crate::upstream_handler::UpstreamType::Udp => {
dns_debug!("开始创建UDP传输: {} ({})", spec.name, spec.server);
let (server, port) = parse_simple_server_address(&spec.server, 53);
dns_debug!("UDP地址解析: server={}, port={}", server, port);
let transport_config = crate::transport::TransportConfig {
server,
port,
timeout: default_timeout,
tcp_fast_open: false,
tcp_nodelay: true,
pool_size: 10,
};
resolver.add_udp_transport(transport_config);
dns_debug!("✅ UDP传输添加成功: {}", spec.name);
},
crate::upstream_handler::UpstreamType::Tcp => {
dns_debug!("开始创建TCP传输: {} ({})", spec.name, spec.server);
let (server, port) = parse_simple_server_address(&spec.server, 53);
dns_debug!("TCP地址解析: server={}, port={}", server, port);
let transport_config = crate::transport::TransportConfig {
server,
port,
timeout: default_timeout,
tcp_fast_open: false,
tcp_nodelay: true,
pool_size: 10,
};
resolver.add_tcp_transport(transport_config);
dns_debug!("✅ TCP传输添加成功: {}", spec.name);
},
crate::upstream_handler::UpstreamType::DoH => {
dns_debug!("开始创建DoH传输: {} ({})", spec.name, spec.server);
if !spec.server.starts_with("https://") {
return Err(DnsError::InvalidConfig("DoH server must be HTTPS URL".to_string()));
}
let (hostname, port) = parse_url_components(&spec.server)?;
dns_debug!("DoH URL解析: hostname={}, port={}", hostname, port);
let connection_server = spec.resolved_ip.as_ref().unwrap_or(&hostname);
dns_debug!("DoH连接服务器: {}", connection_server);
let https_config = crate::transport::HttpsConfig {
base: crate::transport::TransportConfig {
server: connection_server.clone(),
port,
timeout: default_timeout,
tcp_fast_open: false,
tcp_nodelay: true,
pool_size: 5,
},
url: spec.server.clone(),
method: crate::transport::HttpMethod::POST,
user_agent: get_user_agent(),
};
dns_debug!("调用resolver.add_https_transport...");
match resolver.add_https_transport(https_config) {
Ok(_) => dns_debug!("✅ DoH传输添加成功: {}", spec.name),
Err(e) => {
dns_debug!("❌ DoH传输添加失败: {} - 错误: {:?}", spec.name, e);
return Err(e);
}
}
},
crate::upstream_handler::UpstreamType::DoT => {
dns_debug!("开始创建DoT传输: {} ({})", spec.name, spec.server);
let (server, port) = parse_simple_server_address(&spec.server, 853);
dns_debug!("DoT地址解析: server={}, port={}", server, port);
let connection_server = spec.resolved_ip.as_ref().unwrap_or(&server);
dns_debug!("DoT连接服务器: {}, SNI: {}", connection_server, server);
let tls_config = crate::transport::TlsConfig {
base: crate::transport::TransportConfig {
server: connection_server.clone(),
port,
timeout: default_timeout,
tcp_fast_open: false,
tcp_nodelay: true,
pool_size: 5,
},
server_name: server, verify_cert: true,
};
dns_debug!("调用resolver.add_tls_transport...");
match resolver.add_tls_transport(tls_config) {
Ok(_) => dns_debug!("✅ DoT传输添加成功: {}", spec.name),
Err(e) => {
dns_debug!("❌ DoT传输添加失败: {} - 错误: {:?}", spec.name, e);
return Err(e);
}
}
},
}
}
dns_debug!("SmartDnsResolver::new - 所有传输创建完成,解析器构建成功");
Ok(Self {
resolver,
upstream_manager,
decision_engine,
query_strategy,
enable_edns,
})
}
pub async fn query(&self, request: DnsQueryRequest) -> Result<DnsQueryResponse> {
let start_time = Instant::now();
let query_id = request.query_id.clone().unwrap_or_else(|| Uuid::new_v4().to_string());
let result = match self.query_strategy {
QueryStrategy::Fifo => self.query_fifo(&request).await,
QueryStrategy::Smart => self.query_smart(&request).await,
QueryStrategy::RoundRobin => self.query_round_robin(&request).await,
};
let duration = start_time.elapsed();
match result {
Ok((response, server_used)) => {
if let Some(engine) = &self.decision_engine {
engine.update_metrics(&server_used, duration, true, true).await;
}
Ok(DnsQueryResponse {
query_id,
domain: request.domain,
record_type: request.record_type,
success: true,
error: None,
records: self.convert_response_to_records(response),
duration_ms: duration.as_millis() as u64,
server_used: Some(server_used),
dnssec_status: Some(crate::builder::types::DnssecStatus::Indeterminate),
dnssec_records: Vec::new(),
})
},
Err(e) => {
Ok(DnsQueryResponse {
query_id,
domain: request.domain,
record_type: request.record_type,
success: false,
error: Some(format!("查询失败 (策略: {:?}): {}", self.query_strategy, e)),
records: Vec::new(),
duration_ms: duration.as_millis() as u64,
server_used: None,
dnssec_status: Some(crate::builder::types::DnssecStatus::Indeterminate),
dnssec_records: Vec::new(),
})
}
}
}
async fn query_fifo(&self, request: &DnsQueryRequest) -> Result<(crate::Response, String)> {
let record_type = self.convert_record_type(request.record_type);
let client_ip = request.client_address.as_ref()
.and_then(|ip| ip.parse().ok());
if let Some(engine) = &self.decision_engine {
if let Some(spec) = engine.select_fifo_upstream().await {
let start_time = Instant::now();
match self.resolver.query_with_client_ip(&request.domain, record_type, crate::types::QClass::IN, client_ip).await {
Ok(response) => {
let duration = start_time.elapsed();
engine.update_metrics(&spec.name, duration, true, true).await;
Ok((response, spec.name))
},
Err(e) => {
let duration = start_time.elapsed();
engine.update_metrics(&spec.name, duration, false, false).await;
Err(e)
}
}
} else {
Err(DnsError::NoUpstreamAvailable)
}
} else {
Err(DnsError::InvalidConfig("FIFO strategy requires decision engine".to_string()))
}
}
async fn query_smart(&self, request: &DnsQueryRequest) -> Result<(crate::Response, String)> {
let record_type = self.convert_record_type(request.record_type);
let client_ip = request.client_address.as_ref()
.and_then(|ip| ip.parse().ok());
if let Some(engine) = &self.decision_engine {
if let Some(spec) = engine.select_smart_upstream().await {
let start_time = Instant::now();
match self.resolver.query_with_client_ip(&request.domain, record_type, crate::types::QClass::IN, client_ip).await {
Ok(response) => {
let duration = start_time.elapsed();
engine.update_metrics(&spec.name, duration, true, true).await;
Ok((response, spec.name))
},
Err(e) => {
let duration = start_time.elapsed();
engine.update_metrics(&spec.name, duration, false, false).await;
Err(e)
}
}
} else {
Err(DnsError::NoUpstreamAvailable)
}
} else {
Err(DnsError::InvalidConfig("Smart strategy requires decision engine".to_string()))
}
}
fn convert_record_type(&self, record_type: DnsRecordType) -> crate::types::RecordType {
match record_type {
DnsRecordType::A => crate::types::RecordType::A,
DnsRecordType::AAAA => crate::types::RecordType::AAAA,
DnsRecordType::CNAME => crate::types::RecordType::CNAME,
DnsRecordType::MX => crate::types::RecordType::MX,
DnsRecordType::TXT => crate::types::RecordType::TXT,
DnsRecordType::NS => crate::types::RecordType::NS,
DnsRecordType::PTR => crate::types::RecordType::PTR,
DnsRecordType::SRV => crate::types::RecordType::SRV,
DnsRecordType::SOA => crate::types::RecordType::SOA,
DnsRecordType::RRSIG => crate::types::RecordType::Unknown(46), DnsRecordType::DNSKEY => crate::types::RecordType::Unknown(48), DnsRecordType::DS => crate::types::RecordType::Unknown(43), DnsRecordType::NSEC => crate::types::RecordType::Unknown(47), DnsRecordType::NSEC3 => crate::types::RecordType::Unknown(50), }
}
async fn query_round_robin(&self, request: &DnsQueryRequest) -> Result<(crate::Response, String)> {
let record_type = self.convert_record_type(request.record_type);
let client_ip = request.client_address.as_ref()
.and_then(|ip| ip.parse().ok());
if let Some(engine) = &self.decision_engine {
let mut last_error = None;
let mut attempted_servers = Vec::new();
for attempt in 0..3 {
if let Some(spec) = engine.select_round_robin_upstream().await {
attempted_servers.push(spec.name.clone());
let start_time = Instant::now();
match self.resolver.query_with_client_ip(&request.domain, record_type, crate::types::QClass::IN, client_ip).await {
Ok(response) => {
let duration = start_time.elapsed();
engine.update_metrics(&spec.name, duration, true, true).await;
return Ok((response, spec.name));
},
Err(e) => {
let duration = start_time.elapsed();
engine.update_metrics(&spec.name, duration, false, false).await;
last_error = Some(e);
if attempt < 2 {
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
}
}
}
} else {
break;
}
}
if let Some(error) = last_error {
Err(DnsError::Server(format!(
"Round-robin查询失败,已尝试服务器: [{}],最后错误: {}",
attempted_servers.join(", "),
error
)))
} else {
Err(DnsError::NoUpstreamAvailable)
}
} else {
Err(DnsError::InvalidConfig("Round-robin strategy requires decision engine".to_string()))
}
}
fn convert_response_to_records(&self, response: crate::Response) -> Vec<DnsRecord> {
use crate::builder::types::{DnsRecord, DnsRecordValue};
let mut records = Vec::new();
for record in response.answers {
let record_type = match record.rtype {
crate::types::RecordType::A => DnsRecordType::A,
crate::types::RecordType::AAAA => DnsRecordType::AAAA,
crate::types::RecordType::CNAME => DnsRecordType::CNAME,
crate::types::RecordType::MX => DnsRecordType::MX,
crate::types::RecordType::TXT => DnsRecordType::TXT,
crate::types::RecordType::NS => DnsRecordType::NS,
crate::types::RecordType::PTR => DnsRecordType::PTR,
crate::types::RecordType::SRV => DnsRecordType::SRV,
crate::types::RecordType::SOA => DnsRecordType::SOA,
_ => continue,
};
let value = match record.data {
crate::types::RecordData::A(addr) => DnsRecordValue::IpAddr(addr.into()),
crate::types::RecordData::AAAA(addr) => DnsRecordValue::IpAddr(addr.into()),
crate::types::RecordData::CNAME(name) => DnsRecordValue::Domain(name),
crate::types::RecordData::NS(name) => DnsRecordValue::Domain(name),
crate::types::RecordData::PTR(name) => DnsRecordValue::Domain(name),
crate::types::RecordData::TXT(texts) => DnsRecordValue::Text(texts.join(" ")),
crate::types::RecordData::MX { priority, exchange } => {
DnsRecordValue::Mx { priority, exchange }
},
crate::types::RecordData::SRV { priority, weight, port, target } => {
DnsRecordValue::Srv { priority, weight, port, target }
},
_ => continue,
};
records.push(DnsRecord {
name: record.name,
record_type,
value,
ttl: record.ttl,
});
}
records
}
pub async fn get_stats(&self) -> CoreResolverStats {
let mut stats = CoreResolverStats::new(self.query_strategy, self.enable_edns);
if let Some(engine) = &self.decision_engine {
let metrics = engine.get_all_metrics().await;
stats.total_upstreams = metrics.len();
stats.available_upstreams = engine.available_upstream_count().await;
for (name, metric) in metrics {
stats.total_queries += metric.total_queries;
stats.successful_queries += metric.successful_queries;
stats.failed_queries += metric.failed_queries;
if metric.avg_latency < stats.min_latency || stats.min_latency.is_zero() {
stats.min_latency = metric.avg_latency;
stats.fastest_upstream = Some(name.clone());
}
if metric.avg_latency > stats.max_latency {
stats.max_latency = metric.avg_latency;
stats.slowest_upstream = Some(name);
}
}
}
stats.strategy = self.query_strategy;
stats.edns_enabled = self.enable_edns;
stats
}
pub async fn reset_stats(&self) {
if let Some(engine) = &self.decision_engine {
engine.reset_metrics().await;
}
}
pub async fn get_upstream_status(&self) -> Vec<UpstreamStatus> {
let mut status_list = Vec::new();
if let Some(engine) = &self.decision_engine {
let upstreams = engine.get_upstreams().await;
let metrics = engine.get_all_metrics().await;
for upstream in upstreams {
let metric = metrics.get(&upstream.name).cloned().unwrap_or_default();
status_list.push(UpstreamStatus {
name: upstream.name,
server: upstream.server,
transport_type: upstream.transport_type,
is_available: metric.is_available(),
success_rate: metric.success_rate(),
avg_latency: metric.avg_latency,
consecutive_failures: metric.consecutive_failures,
total_queries: metric.total_queries,
last_success: metric.last_success_time,
});
}
}
status_list
}
pub fn query_strategy(&self) -> QueryStrategy {
self.query_strategy
}
pub fn is_edns_enabled(&self) -> bool {
self.enable_edns
}
pub fn get_decision_engine(&self) -> Option<&Arc<SmartDecisionEngine>> {
self.decision_engine.as_ref()
}
async fn check_emergency_status(&self) -> Option<String> {
if let Some(engine) = &self.decision_engine {
if engine.all_upstreams_failed().await {
let emergency_info = engine.get_emergency_response_info().await;
return Some(format!(
"🚨 应急模式激活: {} (策略: {:?})",
emergency_info.emergency_message,
self.query_strategy
));
}
}
None
}
async fn enhance_error_with_emergency_info(&self, original_error: DnsError) -> String {
if let Some(engine) = &self.decision_engine {
let emergency_info = engine.get_emergency_response_info().await;
if emergency_info.all_servers_failed {
format!(
"查询失败 (策略: {:?}): {}\n🚨 应急信息: {}\n📊 失败统计: {}次\n📋 失败服务器: [{}]",
self.query_strategy,
original_error,
emergency_info.emergency_message,
emergency_info.total_failures,
emergency_info.failed_servers.iter()
.map(|s| format!("{} ({}次)", s.name, s.consecutive_failures))
.collect::<Vec<_>>()
.join(", ")
)
} else if emergency_info.total_failures > 0 {
format!(
"查询失败 (策略: {:?}): {}\n⚠️ 部分服务器不可用: {}次失败",
self.query_strategy,
original_error,
emergency_info.total_failures
)
} else {
format!("查询失败 (策略: {:?}): {}", self.query_strategy, original_error)
}
} else {
format!("查询失败 (策略: {:?}, 无决策引擎): {}", self.query_strategy, original_error)
}
}
pub fn upstream_manager(&self) -> &UpstreamManager {
&self.upstream_manager
}
}
impl Clone for SmartDnsResolver {
fn clone(&self) -> Self {
let config = crate::resolver::CoreResolverConfig {
strategy: crate::builder::strategy::QueryStrategy::Smart,
default_timeout: std::time::Duration::from_secs(5),
retry_count: 2,
enable_cache: true,
max_cache_ttl: std::time::Duration::from_secs(3600),
enable_upstream_monitoring: true,
upstream_monitoring_interval: std::time::Duration::from_secs(30),
default_client_address: None,
port: 53,
concurrent_queries: 10,
recursion_desired: true,
buffer_size: 4096,
enable_stats: true,
log_level: rat_logger::LevelFilter::Info,
enable_dns_log_format: true,
};
Self::new(
config,
self.upstream_manager.clone(),
self.decision_engine.clone(),
self.query_strategy,
self.enable_edns,
).expect("Failed to clone SmartDnsResolver")
}
}
#[derive(Debug, Clone)]
pub struct CoreResolverStats {
pub strategy: QueryStrategy,
pub edns_enabled: bool,
pub total_upstreams: usize,
pub available_upstreams: usize,
pub total_queries: u64,
pub successful_queries: u64,
pub failed_queries: u64,
pub min_latency: std::time::Duration,
pub max_latency: std::time::Duration,
pub fastest_upstream: Option<String>,
pub slowest_upstream: Option<String>,
}
impl CoreResolverStats {
pub fn new(strategy: QueryStrategy, edns_enabled: bool) -> Self {
Self {
strategy,
edns_enabled,
total_upstreams: 0,
available_upstreams: 0,
total_queries: 0,
successful_queries: 0,
failed_queries: 0,
min_latency: std::time::Duration::from_millis(0),
max_latency: std::time::Duration::from_millis(0),
fastest_upstream: None,
slowest_upstream: None,
}
}
pub fn success_rate(&self) -> f64 {
if self.total_queries == 0 {
0.0
} else {
self.successful_queries as f64 / self.total_queries as f64
}
}
pub fn avg_latency(&self) -> std::time::Duration {
if self.min_latency.is_zero() && self.max_latency.is_zero() {
std::time::Duration::from_millis(0)
} else {
(self.min_latency + self.max_latency) / 2
}
}
}
#[derive(Debug, Clone)]
pub struct UpstreamStatus {
pub name: String,
pub server: String,
pub transport_type: crate::upstream_handler::UpstreamType,
pub is_available: bool,
pub success_rate: f64,
pub avg_latency: std::time::Duration,
pub consecutive_failures: u32,
pub total_queries: u64,
pub last_success: Option<std::time::Instant>,
}