use crate::{Request, Response, Result, DnsError};
use crate::types::{Query, RecordType, QClass, Flags, ClientAddress};
use crate::transport::{Transport, UdpTransport, TcpTransport, TlsTransport, HttpsTransport};
use crate::transport::{TransportConfig, TlsConfig, HttpsConfig};
use std::fmt::Debug;
use std::sync::Arc;
use std::time::{Duration, Instant};
use std::net::IpAddr;
use tokio::time::timeout;
use std::collections::HashMap;
use crate::{dns_debug, dns_info, dns_error, dns_transport, dns_warn};
pub mod cache;
pub mod health;
use crate::builder::strategy::QueryStrategy;
use cache::DnsCache;
use health::UpstreamMonitor;
#[derive(Debug, Clone)]
pub struct QueryResult {
pub response: Result<Response>,
pub duration: Duration,
pub transport_type: String,
}
#[derive(Debug, Clone)]
pub struct CoreResolver {
transports: Vec<Arc<dyn Transport + Send + Sync + 'static>>,
strategy: QueryStrategy,
cache: Option<Arc<DnsCache>>,
upstream_monitor: Option<Arc<UpstreamMonitor>>,
default_timeout: Duration,
retry_count: usize,
default_client_address: Option<ClientAddress>,
}
#[derive(Debug, Clone)]
pub struct CoreResolverConfig {
pub strategy: QueryStrategy,
pub default_timeout: Duration,
pub retry_count: usize,
pub enable_cache: bool,
pub max_cache_ttl: Duration,
pub enable_upstream_monitoring: bool,
pub upstream_monitoring_interval: Duration,
pub default_client_address: Option<ClientAddress>,
pub port: u16,
pub concurrent_queries: usize,
pub recursion_desired: bool,
pub buffer_size: usize,
pub enable_stats: bool,
pub log_level: rat_logger::LevelFilter,
pub enable_dns_log_format: bool,
}
impl CoreResolverConfig {
pub fn new(
strategy: QueryStrategy,
default_timeout: Duration,
retry_count: usize,
enable_cache: bool,
max_cache_ttl: Duration,
enable_upstream_monitoring: bool,
upstream_monitoring_interval: Duration,
port: u16,
concurrent_queries: usize,
recursion_desired: bool,
buffer_size: usize,
enable_stats: bool,
log_level: rat_logger::LevelFilter,
enable_dns_log_format: bool,
) -> Self {
Self {
strategy,
default_timeout,
retry_count,
enable_cache,
max_cache_ttl,
enable_upstream_monitoring,
upstream_monitoring_interval,
default_client_address: None, port,
concurrent_queries,
recursion_desired,
buffer_size,
enable_stats,
log_level,
enable_dns_log_format,
}
}
}
impl CoreResolver {
pub fn new(config: CoreResolverConfig) -> Self {
let cache = if config.enable_cache {
Some(Arc::new(DnsCache::new(config.max_cache_ttl)))
} else {
None
};
let upstream_monitor = if config.enable_upstream_monitoring {
Some(Arc::new(UpstreamMonitor::with_config(
config.upstream_monitoring_interval,
health::UpstreamConfig {
min_success_rate: 0.7,
max_avg_response_time: std::time::Duration::from_secs(5),
max_consecutive_failures: 3,
recovery_success_count: 2,
stats_window_size: 100,
max_unavailable_duration: std::time::Duration::from_secs(300),
}
)))
} else {
None
};
Self {
transports: Vec::new(),
strategy: config.strategy,
cache,
upstream_monitor,
default_timeout: config.default_timeout,
retry_count: config.retry_count,
default_client_address: config.default_client_address,
}
}
pub fn add_udp_transport(&mut self, config: TransportConfig) {
dns_info!("🪶 添加UDP传输: {}:{}", config.server, config.port);
let transport = Arc::new(UdpTransport::new(config));
self.transports.push(transport.clone());
dns_info!("🪶 UDP传输已添加,当前传输总数: {}", self.transports.len());
dns_debug!("新添加的传输类型: {}", transport.transport_type());
}
pub fn add_tcp_transport(&mut self, config: TransportConfig) {
dns_info!("🔗 添加TCP传输: {}:{}", config.server, config.port);
let transport = Arc::new(TcpTransport::new(config));
self.transports.push(transport.clone());
dns_info!("🔗 TCP传输已添加,当前传输总数: {}", self.transports.len());
dns_debug!("新添加的传输类型: {}", transport.transport_type());
}
pub fn add_tls_transport(&mut self, config: TlsConfig) -> Result<()> {
dns_info!("🔒 添加DoT传输: {}:{}", config.base.server, config.base.port);
let transport = Arc::new(TlsTransport::new(config)?);
self.transports.push(transport.clone());
dns_info!("🔒 DoT传输已添加,当前传输总数: {}", self.transports.len());
dns_debug!("新添加的传输类型: {}", transport.transport_type());
Ok(())
}
pub fn add_https_transport(&mut self, config: HttpsConfig) -> Result<()> {
dns_info!("🌐 添加DoH传输: {}", config.url);
let transport = Arc::new(HttpsTransport::new(config)?);
self.transports.push(transport.clone());
dns_info!("🌐 DoH传输已添加,当前传输总数: {}", self.transports.len());
dns_debug!("新添加的传输类型: {}", transport.transport_type());
Ok(())
}
pub fn add_transport(&mut self, transport: Arc<dyn Transport>) {
self.transports.push(transport);
}
pub async fn query(
&self,
name: &str,
record_type: RecordType,
class: QClass,
) -> Result<Response> {
self.query_with_client_ip(name, record_type, class, None).await
}
pub async fn query_with_client_ip(
&self,
name: &str,
record_type: RecordType,
class: QClass,
client_ip: Option<IpAddr>,
) -> Result<Response> {
let client_address = client_ip.map(|ip| match ip {
IpAddr::V4(addr) => ClientAddress::from_ipv4(addr, 24),
IpAddr::V6(addr) => ClientAddress::from_ipv6(addr, 56),
});
let query = Query {
name: name.to_string(),
qtype: record_type,
qclass: class,
};
if let Some(cache) = &self.cache {
if let Some(cached_response) = cache.get(&query) {
return Ok(cached_response);
}
}
let request = Request {
id: rand::random(),
flags: Flags::default(),
query: query.clone(),
client_address: client_address.or_else(|| self.default_client_address.clone()),
};
let response = self.execute_query_strategy(&request).await?;
if let Some(cache) = &self.cache {
cache.insert(query, response.clone());
}
Ok(response)
}
pub fn set_default_client_address(&mut self, client_address: Option<ClientAddress>) {
self.default_client_address = client_address;
}
pub fn set_default_client_ip(&mut self, client_ip: Option<IpAddr>) {
self.default_client_address = client_ip.map(|ip| match ip {
IpAddr::V4(addr) => ClientAddress::from_ipv4(addr, 24),
IpAddr::V6(addr) => ClientAddress::from_ipv6(addr, 56),
});
}
async fn execute_query_strategy(&self, request: &Request) -> Result<Response> {
if self.transports.is_empty() {
return Err(DnsError::Config("No transports configured".to_string()));
}
dns_info!("🔍 开始DNS查询: {} (类型: {:?}), 策略: {:?}, 可用传输: {}",
request.query.name, request.query.qtype, self.strategy, self.transports.len());
for (i, transport) in self.transports.iter().enumerate() {
dns_debug!("传输[{}]: {}", i, transport.transport_type());
}
match self.strategy {
QueryStrategy::Fifo => self.query_fastest_first(request).await,
QueryStrategy::Smart => self.query_smart_decision(request).await,
QueryStrategy::RoundRobin => self.query_parallel(request).await,
}
}
async fn query_fastest_first(&self, request: &Request) -> Result<Response> {
use tokio::sync::{oneshot, broadcast};
let available_transports = self.get_available_transports();
if available_transports.is_empty() {
return Err(DnsError::Server("No available transports".to_string()));
}
dns_info!("⚡ 使用最快优先策略,并发查询 {} 个传输", available_transports.len());
for (i, transport) in available_transports.iter().enumerate() {
dns_debug!("并发传输[{}]: {}", i, transport.transport_type());
}
let (cancel_tx, _) = broadcast::channel::<()>(1);
let cancel_tx = Arc::new(cancel_tx);
let (success_tx, mut success_rx) = oneshot::channel();
let success_tx = Arc::new(tokio::sync::Mutex::new(Some(success_tx)));
let mut tasks = Vec::new();
for transport in available_transports {
let transport_clone = Arc::clone(&transport);
let request_clone = request.clone();
let mut cancel_rx = cancel_tx.subscribe();
let success_tx_clone = success_tx.clone();
let cancel_tx_clone = cancel_tx.clone();
let upstream_monitor = self.upstream_monitor.clone();
let task = tokio::spawn(async move {
let start = Instant::now();
let transport_type = transport_clone.transport_type();
dns_debug!("🚀 开始使用 {} 传输查询", transport_type);
tokio::select! {
result = transport_clone.send(&request_clone) => {
let duration = start.elapsed();
match result {
Ok(response) => {
dns_info!("✅ {} 传输查询成功 (耗时: {:?}ms)", transport_type, duration.as_millis());
if let Some(upstream_monitor) = &upstream_monitor {
upstream_monitor.record_success(transport_type, duration);
}
if let Ok(mut sender) = success_tx_clone.try_lock() {
if let Some(tx) = sender.take() {
let _ = tx.send(Ok(response));
let _ = cancel_tx_clone.send(());
}
}
}
Err(e) => {
dns_debug!("❌ {} 传输查询失败: {} (耗时: {:?}ms)", transport_type, e, duration.as_millis());
if let Some(upstream_monitor) = &upstream_monitor {
upstream_monitor.record_failure(transport_type);
}
}
}
}
_ = cancel_rx.recv() => {
dns_debug!("传输 {} 的查询任务被取消", transport_clone.transport_type());
}
}
});
tasks.push(task);
}
let (all_done_tx, all_done_rx) = oneshot::channel::<Result<Response>>();
let all_done_tx = Arc::new(tokio::sync::Mutex::new(Some(all_done_tx)));
let all_tasks_handle = tokio::spawn({
let all_done_tx = all_done_tx.clone();
async move {
let _ = futures::future::join_all(tasks).await;
if let Ok(mut sender) = all_done_tx.try_lock() {
if let Some(tx) = sender.take() {
let _ = tx.send(Err(DnsError::Server("All transports failed".to_string())));
}
}
}
});
let result = tokio::select! {
result = &mut success_rx => {
let _ = cancel_tx.send(());
match result {
Ok(response) => response,
Err(_) => Err(DnsError::Server("Internal communication error".to_string()))
}
}
result = all_done_rx => {
result.unwrap_or(Err(DnsError::Server("Internal communication error".to_string())))
}
};
let _ = tokio::time::timeout(Duration::from_millis(100), all_tasks_handle).await;
result
}
async fn query_parallel(&self, request: &Request) -> Result<Response> {
let available_transports = self.get_available_transports();
if available_transports.is_empty() {
return Err(DnsError::Server("No available transports".to_string()));
}
let mut tasks = Vec::new();
for transport in available_transports {
let transport_clone = Arc::clone(&transport);
let request_clone = request.clone();
let task = tokio::spawn(async move {
transport_clone.send(&request_clone).await
});
tasks.push(task);
}
let results = futures::future::join_all(tasks).await;
for result in results {
if let Ok(Ok(response)) = result {
return Ok(response);
}
}
Err(DnsError::Server("All parallel queries failed".to_string()))
}
async fn query_sequential(&self, request: &Request) -> Result<Response> {
let available_transports = self.get_available_transports();
if available_transports.is_empty() {
return Err(DnsError::Server("No available transports".to_string()));
}
let mut last_error = DnsError::Server("No transports tried".to_string());
for transport in available_transports {
for attempt in 0..=self.retry_count {
match transport.send(request).await {
Ok(response) => return Ok(response),
Err(e) => {
last_error = e;
if attempt < self.retry_count {
tokio::time::sleep(Duration::from_millis(100 * (attempt + 1) as u64)).await;
}
}
}
}
}
Err(last_error)
}
async fn query_smart_decision(&self, request: &Request) -> Result<Response> {
let available_transports = self.get_available_transports();
if available_transports.is_empty() {
return Err(DnsError::Server("No available transports".to_string()));
}
let mut tasks = Vec::new();
for (index, transport) in available_transports.iter().enumerate() {
let transport_clone = Arc::clone(transport);
let request_clone = request.clone();
let task = tokio::spawn(async move {
let start = Instant::now();
let transport_type = transport_clone.transport_type();
let result = transport_clone.send(&request_clone).await;
let duration = start.elapsed();
QueryResult {
response: result,
duration,
transport_type: transport_type.to_string(),
}
});
tasks.push(task);
}
let mut results = Vec::new();
let mut fastest_response: Option<Response> = None;
let mut fastest_time = Duration::from_secs(u64::MAX);
let timeout_duration = self.default_timeout;
let deadline = Instant::now() + timeout_duration;
while !tasks.is_empty() && Instant::now() < deadline {
let remaining_time = deadline.duration_since(Instant::now());
match timeout(remaining_time, futures::future::select_all(tasks)).await {
Ok((task_result, _index, remaining_tasks)) => {
tasks = remaining_tasks;
if let Ok(query_result) = task_result {
if let Ok(response) = &query_result.response {
if query_result.duration < fastest_time {
fastest_time = query_result.duration;
fastest_response = Some(response.clone());
}
}
results.push(query_result);
}
}
Err(_) => {
break; }
}
}
let final_result = self.select_best_result(results, fastest_response);
match &final_result {
Ok(response) => {
dns_info!("🧠 Smart策略: 最终选择成功 - 答案数: {}, 查询: {}",
response.answers.len(), request.query.name);
}
Err(e) => {
dns_warn!("🧠 Smart策略: 所有传输均失败 - 错误: {}, 查询: {}",
e, request.query.name);
}
}
final_result
}
fn select_best_result(
&self,
results: Vec<QueryResult>,
fastest_response: Option<Response>,
) -> Result<Response> {
if results.is_empty() {
return Err(DnsError::Timeout);
}
let mut best_response: Option<Response> = None;
let mut best_score = -1i32;
let mut best_duration = Duration::from_secs(u64::MAX);
let mut best_transport_type = String::new();
for result in results.iter() {
if let Ok(response) = &result.response {
let score = response.answers.len() as i32;
if score > best_score ||
(score == best_score && result.duration < best_duration) {
best_score = score;
best_duration = result.duration;
best_response = Some(response.clone());
best_transport_type = result.transport_type.clone();
}
}
}
if let Some(response) = best_response {
dns_info!("🎯 Smart策略: 选择最佳结果 - 传输: {}, 答案数: {}, 耗时: {:?}ms",
best_transport_type, best_score, best_duration.as_millis());
return Ok(response);
}
if let Some(response) = fastest_response {
return Ok(response);
}
for result in results {
if let Err(e) = result.response {
return Err(e);
}
}
Err(DnsError::Server("No valid results".to_string()))
}
fn get_available_transports(&self) -> Vec<Arc<dyn Transport + Send + Sync + 'static>> {
if let Some(upstream_monitor) = &self.upstream_monitor {
self.transports
.iter()
.filter(|t| {
let transport_type = t.transport_type();
upstream_monitor.is_transport_available(transport_type)
})
.cloned()
.collect()
} else {
self.transports.clone()
}
}
pub fn get_transport_stats(&self) -> HashMap<String, (u64, u64, Duration)> {
if let Some(upstream_monitor) = &self.upstream_monitor {
upstream_monitor.get_stats()
} else {
HashMap::new()
}
}
pub fn transport_count(&self) -> usize {
self.transports.len()
}
pub fn clear_cache(&self) {
if let Some(cache) = &self.cache {
cache.clear();
}
}
}