use std::collections::BTreeMap;
use std::convert::{TryFrom, TryInto};
use std::net::SocketAddr;
use std::fmt;
use std::ops::Index;
use std::time::Duration;
use std::time::Instant;
use crate::configuration::DnsConfigReader;
use crate::{DnsConfig, ResolveConfig, error::*};
use crate::internal_error;
use crate::parsers::cfg_resolv_parser::ResolveConfigFamily;
use super::common::*;
bitflags! {
#[derive(Default, Debug, Copy, Clone, PartialEq, Eq)]
pub struct QuerySetupFlags: u16
{
const MEASURE_TIME = 0x0001;
const IGNORE_HOSTS = 0x0002;
const NEXT_NS_ON_OK = 0x0004;
}
}
#[derive(Debug, Clone)]
pub struct QuerySetup
{
pub(crate) flags: QuerySetupFlags,
pub(crate) timeout: Option<Duration>,
}
impl Default for QuerySetup
{
fn default() -> Self
{
return
Self
{
flags:
QuerySetupFlags::empty(),
timeout:
None,
};
}
}
impl QuerySetup
{
pub(crate)
fn get_measure_time(&self) -> bool
{
return self.flags.intersects(QuerySetupFlags::MEASURE_TIME);
}
pub(crate)
fn get_ignore_hosts(&self) -> bool
{
return self.flags.intersects(QuerySetupFlags::IGNORE_HOSTS);
}
pub(crate)
fn get_next_ns_ok(&self) -> bool
{
return self.flags.intersects(QuerySetupFlags::NEXT_NS_ON_OK);
}
pub
fn measure_time(mut self, flag: bool) -> Self
{
self.flags.set(QuerySetupFlags::MEASURE_TIME, flag);
return self;
}
pub
fn ign_hosts(mut self, flag: bool) -> Self
{
self.flags.set(QuerySetupFlags::IGNORE_HOSTS, flag);
return self;
}
pub
fn override_timeout(mut self, timeout: u64) -> Self
{
if timeout == 0
{
return self;
}
self.timeout = Some(Duration::from_secs(timeout));
return self;
}
pub
fn get_timeout(&self, rc: &ResolveConfig) -> Duration
{
return self.timeout.unwrap_or(Duration::from_secs(rc.timeout as u64));
}
pub
fn reset_override_timeout(&mut self)
{
self.timeout = None;
}
pub
fn try_next_ns(mut self, flag: bool) -> Self
{
self.flags.set(QuerySetupFlags::NEXT_NS_ON_OK, flag);
return self;
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum QDnsQueryRec
{
Ok,
ServFail,
NxDomain,
Refused,
NotImpl,
Truncated,
FormError,
}
impl QDnsQueryRec
{
pub(crate)
fn try_next_nameserver(&self, aa: bool, not_aa_retry_next: bool) -> bool
{
match *self
{
Self::Ok =>
{
if not_aa_retry_next == true && aa == false
{
return true;
}
return false;
}
Self::NxDomain =>
{
if aa == true
{
return false;
}
return true;
},
Self::Truncated =>
{
return true;
},
Self::Refused | Self::ServFail | Self::NotImpl | Self::FormError =>
{
return true;
}
}
}
pub(crate)
fn should_try_tcp(&self) -> bool
{
return *self == Self::Truncated || *self == Self::ServFail || *self == Self::NxDomain;
}
}
impl TryFrom<StatusBits> for QDnsQueryRec
{
type Error = CDnsError;
fn try_from(value: StatusBits) -> Result<Self, Self::Error>
{
if value.contains(StatusBits::TRUN_CATION) == true
{
return Ok(QDnsQueryRec::Truncated);
}
else if value.contains(StatusBits::RESP_NOERROR) == true
{
return Ok(QDnsQueryRec::Ok);
}
else if value.contains(StatusBits::RESP_FORMERR) == true
{
return Ok(QDnsQueryRec::FormError);
}
else if value.contains(StatusBits::RESP_NOT_IMPL) == true
{
return Ok(QDnsQueryRec::NotImpl);
}
else if value.contains(StatusBits::RESP_NXDOMAIN) == true
{
return Ok(QDnsQueryRec::NxDomain);
}
else if value.contains(StatusBits::RESP_REFUSED) == true
{
return Ok(QDnsQueryRec::Refused);
}
else if value.contains(StatusBits::RESP_SERVFAIL) == true
{
return Ok(QDnsQueryRec::ServFail);
}
else
{
internal_error!(CDnsErrorType::DnsResponse(CDnsErrorDnsResponse::DnsInvalidStatusBits),
"response status bits unknwon result: '{}'", value.bits());
};
}
}
impl fmt::Display for QDnsQueryRec
{
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result
{
match *self
{
Self::Ok =>
writeln!(f, "OK"),
Self::ServFail =>
writeln!(f, "SERVFAIL"),
Self::NxDomain =>
writeln!(f, "NXDOMAIN"),
Self::Refused =>
writeln!(f, "REFUSED"),
Self::NotImpl =>
writeln!(f, "NOT IMPLEMENTED"),
Self::Truncated =>
writeln!(f, "TRUNCATED"),
Self::FormError =>
writeln!(f, "FORMAT ERROR"),
}
}
}
#[derive(Debug, Default)]
pub struct QDnsQueryResult
{
queries: BTreeMap<u64, CDnsResult<QDnsQuery>>,
}
impl IntoIterator for QDnsQueryResult
{
type Item = (u64, CDnsResult<QDnsQuery>);
type IntoIter = std::collections::btree_map::IntoIter<u64, CDnsResult<QDnsQuery>>;
#[inline]
fn into_iter(self) -> Self::IntoIter
{
return self.queries.into_iter();
}
}
impl Index<u64> for QDnsQueryResult
{
type Output = CDnsResult<QDnsQuery>;
fn index(&self, index: u64) -> &Self::Output
{
return &self.queries.get(&index).unwrap();
}
}
impl QDnsQueryResult
{
pub(crate)
fn new() -> Self
{
return Self{ queries: BTreeMap::new() };
}
pub(crate)
fn push(&mut self, req_id: u64, resp: CDnsResult<QDnsQuery>)
{
self.queries.insert(req_id, resp);
}
pub(crate)
fn push_req(&mut self, req: &DnsRequest, resp: CDnsResult<QDnsQuery>)
{
self.queries.insert(req.get_assigned_id(), resp);
}
pub
fn is_empty(&self) -> bool
{
return self.queries.is_empty();
}
pub
fn contains_dnsreq(&self, req_id: u64) -> bool
{
return self.queries.contains_key(&req_id);
}
pub(crate)
fn extend(&mut self, other: Self, opts: &QuerySetup)
{
for (resp_id, new_resp) in other
{
let Some(prev_resp) = self.queries.get_mut(&resp_id)
else
{
self.queries.insert(resp_id, new_resp);
continue;
};
if new_resp.is_err() == false && prev_resp.is_err() == true
{
*prev_resp = new_resp;
}
else if new_resp.is_err() == true && prev_resp.is_err() == true
{
continue;
}
else if new_resp.as_ref().unwrap().should_check_next_ns(opts.get_next_ns_ok()) == false &&
prev_resp.as_ref().unwrap().should_check_next_ns(opts.get_next_ns_ok()) == true
{
*prev_resp = new_resp;
}
}
}
pub
fn list_results(&self) -> std::collections::btree_map::Iter<'_, u64, Result<QDnsQuery, CDnsError>>
{
return self.queries.iter();
}
pub
fn get_result(self) -> CDnsResult<Vec<QDnsQuery>>
{
let ok = self.collect_ok();
if ok.is_empty() == true
{
internal_error!(CDnsErrorType::DnsNotAvailable, "network error");
}
return Ok(ok);
}
pub
fn get(&self, id: u64) -> Option<&Result<QDnsQuery, CDnsError>>
{
return self.queries.get(&id);
}
pub
fn get_ok_or_error(self) ->CDnsResult<Vec<QDnsQuery>>
{
return
self
.queries
.into_iter()
.map(|e| e.1)
.collect::<CDnsResult<Vec<QDnsQuery>>>();
}
pub
fn collect_ok(self) -> Vec<QDnsQuery>
{
return
self
.queries
.into_iter()
.filter(
|(_k, v)|
v.is_ok()
)
.map(|(_, v)| v.unwrap())
.collect();
}
pub
fn collect_ok_with_id(self) -> BTreeMap<u64, QDnsQuery>
{
return
self
.queries
.into_iter()
.filter(
|(_k, v)|
v.is_ok()
)
.map(|(k, v)| (k, v.unwrap()))
.collect();
}
pub
fn collect_ok_with_answers(self) -> Vec<QDnsQuery>
{
return
self
.queries
.into_iter()
.filter(
|(_k, v)|
v.is_ok() && v.as_ref().unwrap().resp.is_empty() == false
)
.map(|(_, v)| v.unwrap())
.collect();
}
pub
fn collect_ok_with_answers_with_id(self) -> BTreeMap<u64, QDnsQuery>
{
return
self
.queries
.into_iter()
.filter(
|(_k, v)|
v.is_ok() && v.as_ref().unwrap().resp.is_empty() == false
)
.map(|(k, v)| (k, v.unwrap()))
.collect();
}
pub
fn collect_split(self) -> (Vec<(u64, Result<QDnsQuery, CDnsError>)>, Vec<(u64, Result<QDnsQuery, CDnsError>)>)
{
return
self
.queries
.into_iter()
.partition(|(_, v)|
v.is_ok()
);
}
}
impl fmt::Display for QDnsQueryResult
{
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result
{
if self.is_empty() == false
{
for (req, qr) in self.list_results()
{
match qr
{
Ok(r) =>
write!(f, "{}", r)?,
Err(e) =>
write!(f, "request: {}, error: {}", req, e)?
}
}
}
else
{
write!(f, "No DNS server available")?;
}
return Ok(());
}
}
#[derive(Clone, Debug)]
pub struct QDnsQuery
{
pub elapsed: Option<Duration>,
pub server: String,
pub flags: StatusBits,
pub authoratives: Vec<DnsResponsePayload>,
pub resp: Vec<DnsResponsePayload>,
pub status: QDnsQueryRec,
}
impl fmt::Display for QDnsQuery
{
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result
{
write!(f, "Source: {} ", self.server)?;
if let Some(ref el) = self.elapsed
{
write!(f, "{:.2?} ", el)?;
}
if self.is_authorative() == true
{
write!(f, "Authoritative answer")?;
}
else
{
write!(f, "Non-Authoritative answer")?;
}
if self.is_authentic_data() == true
{
write!(f, "AUTHETIC answer")?;
}
else
{
write!(f, "Non-AUTHETIC answer")?;
}
writeln!(f, "\n Authoritatives: {}", self.authoratives.len())?;
if self.authoratives.len() > 0
{
for a in self.authoratives.iter()
{
writeln!(f, "{}", a)?;
}
writeln!(f, "")?;
}
writeln!(f, "Status: {}", self.status)?;
writeln!(f, "Answers: {}", self.resp.len())?;
if self.resp.len() > 0
{
for r in self.resp.iter()
{
writeln!(f, "{}", r)?;
}
writeln!(f, "")?;
}
return Ok(());
}
}
impl IntoIterator for QDnsQuery
{
type Item = DnsResponsePayload;
type IntoIter = std::vec::IntoIter<Self::Item>;
fn into_iter(self) -> Self::IntoIter
{
self.resp.into_iter()
}
}
impl QDnsQuery
{
pub
fn is_ok(&self) -> bool
{
return self.status == QDnsQueryRec::Ok;
}
pub
fn is_authorative(&self) -> bool
{
return self.flags.contains(StatusBits::AUTH_ANSWER);
}
pub
fn get_elapsed_time(&self) -> Option<&Duration>
{
return self.elapsed.as_ref();
}
pub
fn get_server(&self) -> &String
{
return &self.server;
}
pub
fn get_authoratives(&self) -> &[DnsResponsePayload]
{
return self.authoratives.as_slice();
}
pub
fn get_responses(&self) -> &[DnsResponsePayload]
{
return self.resp.as_slice();
}
pub
fn move_responses(self) -> Vec<DnsResponsePayload>
{
return self.resp;
}
pub
fn get_status(&self) -> QDnsQueryRec
{
return self.status;
}
pub(crate)
fn should_check_next_ns(&self, not_aa_retry_next: bool) -> bool
{
return self.status.try_next_nameserver(self.is_authorative(), not_aa_retry_next);
}
pub
fn is_authentic_data(&self) -> bool
{
return self.flags.contains(StatusBits::ANSWER_AUTHN);
}
pub
fn is_checking_disabled(&self) -> bool
{
return self.flags.contains(StatusBits::CHECKING_DISABLED);
}
}
impl QDnsQuery
{
pub(crate)
fn from_local(req_pl: Vec<DnsResponsePayload>, now: Option<&Instant>) -> QDnsQuery
{
let elapsed = now.map(|n| n.elapsed());
return
Self
{
elapsed:
elapsed,
server:
HOST_CFG_PATH.to_string(),
flags:
StatusBits::AUTH_ANSWER | StatusBits::CHECKING_DISABLED | StatusBits::RESP_NOERROR,
authoratives:
Vec::new(),
status:
QDnsQueryRec::Ok,
resp:
req_pl
};
}
pub(crate)
fn from_response(
server: &SocketAddr,
ans_header: DnsHeader,
ans: DnsRequestAnswer,
now: Option<Instant>
) -> CDnsResult<Self>
{
return Ok(
Self
{
elapsed:
now.map_or(None, |n| Some(n.elapsed())),
server:
server.to_string(),
flags:
ans_header.status, authoratives:
ans.authoratives,
status:
ans_header.status.try_into()?,
resp:
ans.response,
}
);
}
}
#[derive(Clone, Debug)]
pub struct QDnsRequests
{
pub(crate) ordered_req_list: Vec<DnsRequest>,
pub(crate) opts: QuerySetup,
}
impl QDnsRequests
{
#[inline]
pub(crate)
fn mirror<'t>(&'t self) -> Vec<&'t DnsRequest>
{
return self.ordered_req_list.iter().collect();
}
}
impl QDnsRequests
{
pub
fn make_empty(opts: QuerySetup) -> Self
{
return
Self
{
ordered_req_list: Vec::new(),
opts: opts,
};
}
pub
fn add_request<R>(&mut self, req_id: u64, qtype: QType, req_name: R) -> CDnsResult<()>
where
R: TryInto<QDnsName, Error = CDnsError>
{
let qnsname = req_name.try_into()?;
self.ordered_req_list.push(DnsRequest::construct_lookup(req_id, qnsname, qtype)?);
return Ok(());
}
pub
fn resolve_fqdn<R>(cfg: &DnsConfig<ResolveConfig>, req_id_4: u64, req_id_6: u64, req_name: R, opts: QuerySetup) -> CDnsResult<Self>
where
R: TryInto<QDnsName, Error = CDnsError>
{
return
Self::resolve_a_aaaa_request(req_id_4, req_id_6,
cfg.read_config().family, req_name, opts);
}
pub
fn resolve_a_aaaa_request<R>(req4_id: u64, req6_id: u64, res_fam: ResolveConfigFamily,
req_name_ref: R, opts: QuerySetup
) -> CDnsResult<Self>
where
R: TryInto<QDnsName, Error = CDnsError>
{
let req_n = req_name_ref.try_into()?;
let reqs =
match res_fam
{
ResolveConfigFamily::INET4_INET6 =>
{
vec![
DnsRequest::construct_lookup(req4_id, req_n.clone(), QType::A)?,
DnsRequest::construct_lookup(req6_id, req_n, QType::AAAA)?
]
},
ResolveConfigFamily::INET6_INET4 =>
{
vec![
DnsRequest::construct_lookup(req6_id, req_n.clone(), QType::AAAA)?,
DnsRequest::construct_lookup(req4_id, req_n, QType::A)?
]
},
ResolveConfigFamily::INET6 =>
{
vec![
DnsRequest::construct_lookup(req6_id, req_n, QType::AAAA)?,
]
},
ResolveConfigFamily::INET4 =>
{
vec![
DnsRequest::construct_lookup(req4_id, req_n.clone(), QType::A)?
]
}
_ =>
{
vec![
DnsRequest::construct_lookup(req4_id, req_n.clone(), QType::A)?,
DnsRequest::construct_lookup(req6_id, req_n, QType::AAAA)?,
]
}
};
return Ok(
Self
{
ordered_req_list: reqs,
opts: opts,
}
);
}
pub
fn resolve_reverse<R>(req_id: u64, fqdn: R, opts: QuerySetup) -> CDnsResult<QDnsRequests>
where
R: AsRef<str>
{
let mut dns_req = QDnsRequests::make_empty(opts);
dns_req.add_request(req_id, QType::PTR, fqdn.as_ref())?;
return Ok(dns_req);
}
pub
fn resolve_soa<R>(req_id: u64, fqdn: R, opts: QuerySetup) -> CDnsResult<QDnsRequests>
where
R: AsRef<str>
{
let mut dns_req = QDnsRequests::make_empty(opts);
dns_req.add_request(req_id, QType::SOA, fqdn.as_ref())?;
return Ok(dns_req);
}
pub
fn resolve_mx<R>(req_id: u64, fqdn: R, opts: QuerySetup) -> CDnsResult<QDnsRequests>
where
R: AsRef<str>
{
let mut dns_req = QDnsRequests::make_empty(opts);
dns_req.add_request(req_id, QType::MX, fqdn.as_ref())?;
return Ok(dns_req);
}
}