use std::
{
collections::BTreeMap, io::{self, ErrorKind}, net::{IpAddr, SocketAddr, TcpStream, UdpSocket}, rc::Rc, sync::Arc, time::Instant, vec
};
use instance_copy_on_write::{ICoWWeak, ICoWWeakRead};
use crate::{CDnsError, common::{CDdnsGlobals, DnsRequestAnswerParser, QDnsName}};
#[cfg(feature = "use_sync_tls")]
use crate::network::with_tls::{TcpHttpsConnection, TcpTlsConnection};
use crate::network::{NetworkTap, SocketTap};
use crate::parsers::cfg_resolv_parser::{ConfigEntryTls, ResolveConfEntry};
use crate::
{
CDnsErrorDnsResponse,
CDnsErrorType,
CDnsIOError,
DnsResponsePayload,
ErrorReport,
GLOBAL_CONFIG,
HostConfig,
QDnsQuery,
ResolveConfig,
ResolverStdErr,
internal_error_map,
CDnsResult,
QDnsQueryResult,
QType,
QuerySetup,
common::DnsRequest,
configuration::{ConfigGetter, DnsConfigReader},
query::QDnsRequests
};
#[allow(async_fn_in_trait)]
pub trait ToSocketAddrs
{
type Iter: Iterator<Item = SocketAddr>;
async fn to_socket_addrs(&self) -> io::Result<Self::Iter>;
}
#[derive(Clone, Debug)]
pub enum QDnsSockerAddr
{
Ip(SocketAddr),
Host(String, u16)
}
impl ToSocketAddrs for QDnsSockerAddr
{
type Iter = vec::IntoIter<SocketAddr>;
async
fn to_socket_addrs(&self) -> std::io::Result<Self::Iter>
{
match self
{
QDnsSockerAddr::Ip(socket_addr) =>
return Ok(vec![socket_addr.clone()].into_iter()),
QDnsSockerAddr::Host(host, port) =>
{
let cfg =
GLOBAL_CONFIG
.get()
.ok_or(
std::io::Error::new(
ErrorKind::NotSeekable,
CDnsIOError::from(
internal_error_map!(CDnsErrorType::ConfigError,
"global config is not available")
)
)
)?;
let req0 =
QDnsRequests
::resolve_fqdn(&cfg.get_resolv_l(), 1, 2, host,
QuerySetup::default())
.map_err(|e|
std::io::Error::new(ErrorKind::InvalidData,
CDnsIOError::from(e))
)?;
let mut resolver_res =
Resolver::<ResolverStdErr>::new(cfg, ResolverStdErr);
let res =
resolver_res
.query(&req0)
.await
.map_err(|e|
std::io::Error::new(ErrorKind::Other, CDnsIOError::from(e)))?
.collect_ok()
.into_iter()
.map(|resp|
resp
.get_responses()
.iter()
.map(|r|
r.get_rdata().get_ip().map(|f| SocketAddr::new(f, *port))
)
.collect::<Vec<Option<SocketAddr>>>()
)
.flatten()
.filter(|p| p.is_some())
.map(|v| v.unwrap())
.collect::<Vec<SocketAddr>>();
return Ok( res.into_iter() );
}
}
}
}
impl QDnsSockerAddr
{
pub
fn resolve<D>(host: D) -> std::io::Result<Self>
where D: AsRef<str>
{
if let Ok(addr) = host.as_ref().parse::<SocketAddr>()
{
return Ok(Self::Ip(addr));
}
else
{
let ref_host = host.as_ref();
let (domain, port) =
match ref_host.split_once(":")
{
Some((h, portno)) =>
{
let port: u16 =
portno
.parse()
.map_err(|e|
std::io::Error::new(ErrorKind::InvalidData, format!("{}", e))
)?;
(h, port)
},
None =>
return Err(std::io::Error::new(ErrorKind::InvalidData, "missing prot number"))
};
return Ok(Self::Host(domain.to_string(), port));
}
}
pub
fn resolve_port<D>(host: D, port: u16) -> std::io::Result<Self>
where D: AsRef<str>
{
if let Ok(addr) = host.as_ref().parse::<IpAddr>()
{
return Ok(Self::Ip(SocketAddr::new(addr, port)));
}
else
{
return Ok(Self::Host(host.as_ref().to_string(), port));
}
}
}
#[derive(Debug)]
struct PersistantConnections
{
persistant_sockets: Vec<Rc<Box<dyn SocketTap>>>,
}
impl PersistantConnections
{
fn new() -> Self
{
return Self{ persistant_sockets: Vec::new() };
}
fn invalidate(&mut self)
{
self.persistant_sockets.clear();
}
fn get(&self, rce: &ResolveConfEntry, force_reopen: bool) -> Option<Rc<Box<dyn SocketTap>>>
{
for item in self.persistant_sockets.iter()
{
if item.as_ref().as_ref() == rce
{
if force_reopen == true
{
if let Err(_) = item.reconnect()
{
return None;
}
}
return Some(item.clone());
}
}
return None;
}
fn add(&mut self, sock_tap: Box<dyn SocketTap>) -> Rc<Box<dyn SocketTap>>
{
let st = Rc::new(sock_tap);
self.persistant_sockets.push(st.clone());
return st;
}
}
#[derive(Debug)]
pub struct Resolver<ERR: ErrorReport>
{
resolv_cfg: ICoWWeak<ResolveConfig>,
hosts_cfg: ICoWWeak<HostConfig>,
persistant_sockets: PersistantConnections,
err: ERR,
}
impl<ERR: ErrorReport> Resolver<ERR>
{
#[inline]
fn create_socket(
&mut self,
opts: &QuerySetup,
force_tcp: bool,
resolver: &Arc<ResolveConfEntry>,
resolv_cfg: &ResolveConfig
) -> CDnsResult<Rc<Box<dyn SocketTap>>>
{
let is_tls = resolver.get_tls_type();
if is_tls == ConfigEntryTls::Tls
{
#[cfg(feature = "use_sync_tls")]
{
let Some(socket) =
self.persistant_sockets.get(resolver.as_ref(), resolv_cfg.option_flags.is_reopen_socket())
else
{
let socket =
NetworkTap
::<TcpTlsConnection>
::new_tls(resolver.clone(), opts.get_timeout(resolv_cfg), CDdnsGlobals::get_tcp_conn_timeout())?;
return Ok(self.persistant_sockets.add(socket));
};
return Ok(socket);
}
#[cfg(not(feature = "use_sync_tls"))]
internal_error!(CDnsErrorType::SocketNotSupported,
"socket not supported: '{}'", resolver.get_tls_type());
}
else if is_tls == ConfigEntryTls::Https
{
#[cfg(feature = "use_sync_tls")]
{
let Some(socket) =
self.persistant_sockets.get(resolver.as_ref(), resolv_cfg.option_flags.is_reopen_socket())
else
{
let socket =
NetworkTap
::<TcpHttpsConnection>
::new_https(resolver.clone(), opts.get_timeout(resolv_cfg), CDdnsGlobals::get_tcp_conn_timeout())?;
return Ok(self.persistant_sockets.add(socket));
};
return Ok(socket);
}
#[cfg(not(feature = "use_sync_tls"))]
internal_error!(CDnsErrorType::SocketNotSupported,
"socket not supported: '{}'", resolver.get_tls_type());
}
else if resolv_cfg.option_flags.is_force_tcp() == true || force_tcp == true
{
let Some(socket) =
self.persistant_sockets.get(resolver.as_ref(), resolv_cfg.option_flags.is_reopen_socket())
else
{
let socket =
NetworkTap
::<TcpStream>
::new_tcp(resolver.clone(), opts.get_timeout(resolv_cfg),
CDdnsGlobals::get_tcp_conn_timeout())?;
return Ok(self.persistant_sockets.add(socket));
};
return Ok(socket);
}
else
{
let socket =
NetworkTap
::<UdpSocket>
::new_udp(resolver.clone(), opts.get_timeout(resolv_cfg))
.map(|s| Rc::new(s));
return socket;
}
}
fn lookup_hosts_internal(&self, hlist: &HostConfig, qname: &QDnsName, qtype: QType, opts: &QuerySetup) -> CDnsResult<Option<QDnsQuery>>
{
if qtype.is_hosts() == false
{
return Ok(None);
}
let now =
if opts.get_measure_time() == true
{
Some(Instant::now())
}
else
{
None
};
if let Some(ip) = <&QDnsName as Into<Option<IpAddr>>>::into(qname)
{
let Some(host_name_ent) = hlist.search_by_ip(&ip)
else { return Ok(None) };
let drp = DnsResponsePayload::new_local(qtype, host_name_ent).unwrap();
return Ok(Some(QDnsQuery::from_local(drp, now.as_ref())));
}
else
{
let req_name = qname.get_query();
let Some(host_name_ent) = hlist.search_by_fqdn(qtype, req_name)
else { return Ok(None) };
let drp = DnsResponsePayload::new_local(qtype, host_name_ent).unwrap();
return Ok(Some(QDnsQuery::from_local(drp, now.as_ref())));
}
}
async
fn process_hosts(&self, reqs: &mut Vec<&DnsRequest>, opts: &QuerySetup) -> CDnsResult<QDnsQueryResult>
{
let mut qqr = QDnsQueryResult::default();
if opts.get_ignore_hosts() == false
{
let hlist =
self.hosts_cfg.aquire()
.ok_or(
internal_error_map!(CDnsErrorType::ConfigError, "resolver's config has gone!")
)?;
reqs.retain(
|r|
{
if r.payload.qtype.is_hosts() == false
{
return true;
}
match self.lookup_hosts_internal(&hlist, &r.get_name(), r.get_type(), opts)
{
Ok(Some(resp)) =>
{
qqr.push(r.get_assigned_id(), Ok(resp));
return false;
},
Ok(None) =>
{
return true;
},
Err(e) =>
{
self.err.report(e);
return true;
}
}
}
);
}
return Ok(qqr);
}
fn process_request(
&mut self,
resolv_cfg: &ICoWWeakRead<'_, ResolveConfig>,
reqs: &mut Vec<&DnsRequest>,
opts: &QuerySetup,
) -> QDnsQueryResult
{
let mut responses: QDnsQueryResult = QDnsQueryResult::new();
if resolv_cfg.item_updated() == true
{
self.persistant_sockets.invalidate();
}
for _ in 0..resolv_cfg.get_attempts()
{
for resolver in resolv_cfg.get_resolvers_iter()
{
let res =
if resolv_cfg.option_flags.is_no_parallel() == true
{
self.query_exec_seq(opts, reqs, resolver, &resolv_cfg, false)
}
else
{
self.query_exec_pipelined(opts, reqs, resolver, &resolv_cfg, false)
};
match res
{
Ok(resp) =>
{
responses.extend(resp, opts);
},
Err(e) =>
{
self.err.report(e);
}
}
if reqs.is_empty() == true
{
break;
}
}
}
return responses;
}
fn query_exec_pipelined(
&mut self,
opts: &QuerySetup,
reqs: &mut Vec<&DnsRequest>,
resolver: &Arc<ResolveConfEntry>,
resolv_cfg: &ICoWWeakRead<'_, ResolveConfig>,
requery: bool,
) -> CDnsResult<QDnsQueryResult>
{
let force_tcp = resolv_cfg.option_flags.is_force_tcp() | requery;
let tap =
self.create_socket(opts, force_tcp, &resolver, resolv_cfg)?;
let mut req_pkt_map: BTreeMap<u16, (&DnsRequest, Option<Instant>)> = BTreeMap::new();
while let Some(req) = reqs.pop()
{
let rand_id =
loop
{
let rand_id = rand::random();
if req_pkt_map.contains_key(&rand_id) == false
{
break rand_id;
}
};
let pkt = req.to_bytes(tap.should_append_len(), rand_id)?;
let now =
if opts.get_measure_time() == true
{
Some(Instant::now())
}
else
{
None
};
tap.send(pkt.as_slice())?;
req_pkt_map.insert(rand_id, (req, now));
}
let mut responses = QDnsQueryResult::default();
let mut tcp_requery: Vec<&DnsRequest> = Vec::new();
while req_pkt_map.is_empty() == false
{
let Some(ans_pkt) = tap.recv()?
else
{
self.err.report(
internal_error_map!(CDnsErrorType::RequestTimeout,
"timeout waiting for responses count: '{}'", req_pkt_map.len())
);
for (_, (req, _)) in req_pkt_map
{
responses.push(
req.get_assigned_id(),
Err(
internal_error_map!(CDnsErrorType::RequestTimeout,
"request timeout for {}", req)
)
);
reqs.push(req);
}
break;
};
let req_ans_par = DnsRequestAnswerParser::try_from(ans_pkt.as_slice())?;
let req_ans_id = req_ans_par.get_pkt_id();
let Some((mapped_req, inst_time)) = req_pkt_map.remove(&req_ans_id)
else
{
self.err.report(
internal_error_map!(CDnsErrorType::DnsResponse(CDnsErrorDnsResponse::DnsUnknownReqID),
"received pkt with unknown req id: {}", req_ans_id)
);
continue;
};
let mut ans = req_ans_par.into_answer(mapped_req.get_assigned_id())?;
ans.1.verify(None, &ans.0, mapped_req)?;
let res_query =
QDnsQuery::from_response(tap.get_remote_addr(), ans.0, ans.1, inst_time);
if let Ok(ref qdns) = res_query
{
if qdns.get_status().should_try_tcp() == true && force_tcp == false
{
tcp_requery.push(mapped_req);
}
else if qdns.should_check_next_ns(opts.get_next_ns_ok()) == true
{
responses.push(mapped_req.get_assigned_id(), res_query);
reqs.push(mapped_req);
}
else
{
responses.push(mapped_req.get_assigned_id(), res_query);
}
}
else
{
responses.push(mapped_req.get_assigned_id(), res_query);
}
}
if tcp_requery.is_empty() == false
{
let res =
self.query_exec_pipelined(opts, &mut tcp_requery, resolver, resolv_cfg, true)?;
responses.extend(res, opts);
if tcp_requery.is_empty() == false
{
reqs.extend(tcp_requery);
}
}
return Ok(responses);
}
fn query_exec_seq(
&mut self,
opts: &QuerySetup,
reqs: &mut Vec<&DnsRequest>,
resolver: &Arc<ResolveConfEntry>,
resolv_cfg: &ICoWWeakRead<'_, ResolveConfig>,
requery: bool,
) -> CDnsResult<QDnsQueryResult>
{
let mut responses = QDnsQueryResult::default();
let force_tcp = resolv_cfg.option_flags.is_force_tcp() || requery;
let mut tcp_requery: Vec<&DnsRequest> = Vec::new();
let mut next_ns_req: Vec<&DnsRequest> = Vec::new();
let tap =
self
.create_socket(opts, force_tcp, &resolver, resolv_cfg)?;
while let Some(req) = reqs.pop()
{
let pkt_ran_id = rand::random();
let pkt = req.to_bytes(tap.should_append_len(), pkt_ran_id)?;
let now =
if opts.get_measure_time() == true
{
Some(Instant::now())
}
else
{
None
};
tap.send(pkt.as_slice())?;
let Some(ans_pkt) = tap.recv()?
else
{
self.err.report(
internal_error_map!(CDnsErrorType::RequestTimeout,
"timeout waiting for request: '{}'", req)
);
return Err(
internal_error_map!(CDnsErrorType::RequestTimeout, "request timeout {}", req)
);
};
let mut ans =
DnsRequestAnswerParser::try_from(ans_pkt.as_slice())?
.into_answer(req.get_assigned_id())?;
ans.1.verify(Some(pkt_ran_id), &ans.0, req)?;
let res_query = QDnsQuery::from_response(tap.get_remote_addr(), ans.0, ans.1, now);
if let Ok(ref qdns) = res_query
{
if qdns.get_status().should_try_tcp() == true && force_tcp == false
{
tcp_requery.push(req);
}
else if qdns.should_check_next_ns(opts.get_next_ns_ok()) == true
{
responses.push_req(req, res_query);
next_ns_req.push(req);
}
else
{
responses.push_req(req, res_query);
}
}
else
{
responses.push_req(req, res_query);
}
}
if tcp_requery.is_empty() == false
{
let res =
self.query_exec_seq(opts, &mut tcp_requery, resolver, resolv_cfg, true)?;
responses.extend(res, opts);
if tcp_requery.is_empty() == false
{
reqs.extend(tcp_requery);
}
}
reqs.extend(next_ns_req);
return Ok(responses);
}
}
impl<ERR: ErrorReport> Resolver<ERR>
{
pub
fn new<CFG: ConfigGetter>(config: CFG, err: ERR) -> Self
{
return
Self
{
resolv_cfg:
config.get_resolv().get_reader(),
hosts_cfg:
config.get_hosts().get_reader(),
persistant_sockets:
PersistantConnections::new(),
err:
err,
};
}
pub
fn lookup_hosts<R>(&self, req: R, qtype: QType, opts: &QuerySetup) -> CDnsResult<Option<QDnsQuery>>
where
R: TryInto<QDnsName, Error = CDnsError>
{
let hlist =
self.hosts_cfg.aquire()
.ok_or(
internal_error_map!(CDnsErrorType::ConfigError, "resolver's config has gone!")
)?;
return self.lookup_hosts_internal(&hlist, &req.try_into()?, qtype, opts);
}
pub async
fn query(&mut self, reqs: &QDnsRequests) -> CDnsResult<QDnsQueryResult>
{
let resolv_cfg =
self.resolv_cfg.aquire()
.ok_or(
internal_error_map!(CDnsErrorType::ConfigError, "resolver's config has gone!")
)?;
let mut mirrowed_reqs = reqs.mirror();
let qqr =
if resolv_cfg.lookup.is_file_first() == true
{
let mut qqr = self.process_hosts(&mut mirrowed_reqs, &reqs.opts)?;
if mirrowed_reqs.is_empty() == false
{
let pr = self.process_request(&resolv_cfg, &mut mirrowed_reqs, &reqs.opts);
qqr.extend(pr, &reqs.opts);
}
qqr
}
else
{
let mut qqr = self.process_request(&resolv_cfg, &mut mirrowed_reqs, &reqs.opts);
if resolv_cfg.lookup.is_file_first() == false && mirrowed_reqs.is_empty() == false
{
let res = self.process_hosts(&mut mirrowed_reqs, &reqs.opts)?;
qqr.extend(res, &reqs.opts);
}
qqr
};
return Ok(qqr);
}
}