use crate::{
transport::{Transport, TransportConfig, HttpsConfig, TlsConfig},
utils::{parse_server_address, parse_url_components, get_user_agent},
Result, DnsError,
dns_info, dns_debug,
};
use std::{
collections::HashMap,
time::Duration,
};
use async_trait::async_trait;
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub enum UpstreamType {
Udp,
Tcp,
DoT,
DoH,
}
#[derive(Debug, Clone)]
pub struct UpstreamSpec {
pub name: String,
pub transport_type: UpstreamType,
pub server: String,
pub resolved_ip: Option<String>,
pub weight: u32,
pub region: Option<String>,
}
#[async_trait]
pub trait UpstreamHandler: Send + Sync + std::fmt::Debug {
fn handler_type(&self) -> UpstreamType;
async fn create_transport(&self, spec: &UpstreamSpec) -> Result<Box<dyn Transport>>;
fn validate_spec(&self, spec: &UpstreamSpec) -> Result<()>;
fn default_port(&self) -> u16;
}
#[derive(Debug, Default)]
pub struct UdpHandler;
#[async_trait]
impl UpstreamHandler for UdpHandler {
fn handler_type(&self) -> UpstreamType {
UpstreamType::Udp
}
async fn create_transport(&self, spec: &UpstreamSpec) -> Result<Box<dyn Transport>> {
let (server, port) = parse_server_address(&spec.server, self.default_port())?;
let actual_server = spec.resolved_ip.as_ref().unwrap_or(&server);
let config = TransportConfig {
server: actual_server.clone(),
port,
timeout: Duration::from_secs(5),
tcp_fast_open: false,
tcp_nodelay: true,
pool_size: 10,
};
Ok(Box::new(crate::transport::UdpTransport::new(config)))
}
fn validate_spec(&self, spec: &UpstreamSpec) -> Result<()> {
if spec.server.is_empty() {
return Err(DnsError::InvalidConfig("UDP server cannot be empty".to_string()));
}
Ok(())
}
fn default_port(&self) -> u16 {
53
}
}
#[derive(Debug, Default)]
pub struct TcpHandler;
#[async_trait]
impl UpstreamHandler for TcpHandler {
fn handler_type(&self) -> UpstreamType {
UpstreamType::Tcp
}
async fn create_transport(&self, spec: &UpstreamSpec) -> Result<Box<dyn Transport>> {
let (server, port) = parse_server_address(&spec.server, self.default_port())?;
let actual_server = spec.resolved_ip.as_ref().unwrap_or(&server);
let config = TransportConfig {
server: actual_server.clone(),
port,
timeout: Duration::from_secs(5),
tcp_fast_open: false,
tcp_nodelay: true,
pool_size: 10,
};
Ok(Box::new(crate::transport::TcpTransport::new(config)))
}
fn validate_spec(&self, spec: &UpstreamSpec) -> Result<()> {
if spec.server.is_empty() {
return Err(DnsError::InvalidConfig("TCP server cannot be empty".to_string()));
}
Ok(())
}
fn default_port(&self) -> u16 {
53
}
}
#[derive(Debug, Default)]
pub struct DoTHandler;
#[async_trait]
impl UpstreamHandler for DoTHandler {
fn handler_type(&self) -> UpstreamType {
UpstreamType::DoT
}
async fn create_transport(&self, spec: &UpstreamSpec) -> Result<Box<dyn Transport>> {
let (server, port) = parse_server_address(&spec.server, self.default_port())?;
let connection_server = spec.resolved_ip.as_ref().unwrap_or(&server);
let sni_name = server.clone();
let config = TlsConfig {
base: TransportConfig {
server: connection_server.clone(),
port,
timeout: Duration::from_secs(10),
tcp_fast_open: false,
tcp_nodelay: true,
pool_size: 5,
},
server_name: sni_name,
verify_cert: true,
};
Ok(Box::new(crate::transport::TlsTransport::new(config)?))
}
fn validate_spec(&self, spec: &UpstreamSpec) -> Result<()> {
if spec.server.is_empty() {
return Err(DnsError::InvalidConfig("DoT server cannot be empty".to_string()));
}
Ok(())
}
fn default_port(&self) -> u16 {
853
}
}
#[derive(Debug, Default)]
pub struct DoHHandler;
#[async_trait]
impl UpstreamHandler for DoHHandler {
fn handler_type(&self) -> UpstreamType {
UpstreamType::DoH
}
async fn create_transport(&self, spec: &UpstreamSpec) -> Result<Box<dyn Transport>> {
let url = &spec.server;
let (hostname, port) = parse_url_components(url)?;
let connection_server = spec.resolved_ip.as_ref().unwrap_or(&hostname);
let config = HttpsConfig {
base: TransportConfig {
server: connection_server.clone(),
port,
timeout: Duration::from_secs(10),
tcp_fast_open: false,
tcp_nodelay: true,
pool_size: 5,
},
url: url.clone(),
method: crate::transport::HttpMethod::POST,
user_agent: get_user_agent(),
};
Ok(Box::new(crate::transport::HttpsTransport::new(config)?))
}
fn validate_spec(&self, spec: &UpstreamSpec) -> Result<()> {
if spec.server.is_empty() {
return Err(DnsError::InvalidConfig("DoH server cannot be empty".to_string()));
}
if !spec.server.starts_with("https://") {
return Err(DnsError::InvalidConfig("DoH URL must use HTTPS".to_string()));
}
Ok(())
}
fn default_port(&self) -> u16 {
443
}
}
#[derive(Debug)]
pub struct UpstreamManager {
handlers: HashMap<UpstreamType, Box<dyn UpstreamHandler>>,
specs: Vec<UpstreamSpec>,
}
impl Clone for UpstreamManager {
fn clone(&self) -> Self {
let mut new_manager = Self::default();
new_manager.specs = self.specs.clone();
new_manager
}
}
impl Default for UpstreamManager {
fn default() -> Self {
let mut handlers: HashMap<UpstreamType, Box<dyn UpstreamHandler>> = HashMap::new();
handlers.insert(UpstreamType::Udp, Box::new(UdpHandler));
handlers.insert(UpstreamType::Tcp, Box::new(TcpHandler));
handlers.insert(UpstreamType::DoT, Box::new(DoTHandler));
handlers.insert(UpstreamType::DoH, Box::new(DoHHandler));
Self {
handlers,
specs: Vec::new(),
}
}
}
impl UpstreamManager {
pub fn new() -> Self {
Self::default()
}
pub fn add_upstream(&mut self, spec: UpstreamSpec) -> Result<()> {
dns_info!("Adding upstream server: {} ({:?}) -> {}", spec.name, spec.transport_type, spec.server);
if let Some(handler) = self.handlers.get(&spec.transport_type) {
handler.validate_spec(&spec)?;
} else {
return Err(DnsError::InvalidConfig(
format!("Unsupported transport type: {:?}", spec.transport_type)
));
}
self.specs.push(spec);
dns_debug!("Successfully added upstream server, total count: {}", self.specs.len());
Ok(())
}
pub async fn create_transport(&self, spec: &UpstreamSpec) -> Result<Box<dyn Transport>> {
if let Some(handler) = self.handlers.get(&spec.transport_type) {
handler.create_transport(spec).await
} else {
Err(DnsError::InvalidConfig(
format!("No handler for transport type: {:?}", spec.transport_type)
))
}
}
pub fn get_specs(&self) -> &[UpstreamSpec] {
&self.specs
}
pub fn filter_by_type(&self, transport_type: UpstreamType) -> Vec<&UpstreamSpec> {
self.specs.iter()
.filter(|spec| spec.transport_type == transport_type)
.collect()
}
}
impl UpstreamSpec {
pub fn udp(name: String, server: String) -> Self {
Self {
name,
transport_type: UpstreamType::Udp,
server,
resolved_ip: None,
weight: 1,
region: None,
}
}
pub fn tcp(name: String, server: String) -> Self {
Self {
name,
transport_type: UpstreamType::Tcp,
server,
resolved_ip: None,
weight: 1,
region: None,
}
}
pub fn dot(name: String, server: String) -> Self {
Self {
name,
transport_type: UpstreamType::DoT,
server,
resolved_ip: None,
weight: 1,
region: None,
}
}
pub fn doh(name: String, url: String) -> Self {
Self {
name,
transport_type: UpstreamType::DoH,
server: url,
resolved_ip: None,
weight: 1,
region: None,
}
}
pub fn with_resolved_ip(mut self, ip: String) -> Self {
self.resolved_ip = Some(ip);
self
}
pub fn with_weight(mut self, weight: u32) -> Self {
self.weight = weight;
self
}
pub fn with_region(mut self, region: String) -> Self {
self.region = Some(region);
self
}
}