use std::{io, ops};
use std::future::Future;
use std::net::IpAddr;
use std::pin::Pin;
use bytes::Bytes;
use futures::future::FutureExt;
#[cfg(feature = "sync")] use tokio::runtime;
use domain::base::iana::Rcode;
use domain::base::message::Message;
use domain::base::message_builder::{
AdditionalBuilder, MessageBuilder, StreamTarget
};
use domain::base::name::{ToDname, ToRelativeDname};
use domain::base::octets::Octets512;
use domain::base::question::Question;
use super::conf::{ResolvConf, ResolvOptions, SearchSuffix};
use super::net::{ServerInfo, ServerList, ServerListCounter};
use crate::lookup::addr::{lookup_addr, FoundAddrs};
use crate::lookup::host::{lookup_host, search_host, FoundHosts};
use crate::lookup::srv::{lookup_srv, FoundSrvs, SrvError};
use crate::resolver::{Resolver, SearchNames};
#[derive(Clone, Debug)]
pub struct StubResolver {
preferred: ServerList,
stream: ServerList,
options: ResolvOptions,
}
impl StubResolver {
pub fn new() -> Self {
Self::from_conf(ResolvConf::default())
}
pub fn from_conf(conf: ResolvConf) -> Self {
StubResolver {
preferred: ServerList::from_conf(&conf, |s| {
s.transport.is_preferred()
}),
stream: ServerList::from_conf(&conf, |s| {
s.transport.is_stream()
}),
options: conf.options
}
}
pub fn options(&self) -> &ResolvOptions {
&self.options
}
pub async fn query<N: ToDname, Q: Into<Question<N>>>(
&self, question: Q
) -> Result<Answer, io::Error> {
Query::new(self)?.run(
Query::create_message(question.into())
).await
}
async fn query_message(
&self, message: QueryMessage
) -> Result<Answer, io::Error> {
Query::new(self)?.run(message).await
}
}
impl StubResolver {
pub async fn lookup_addr(
&self, addr: IpAddr
) -> Result<FoundAddrs<&Self>, io::Error> {
lookup_addr(&self, addr).await
}
pub async fn lookup_host(
&self, qname: impl ToDname
) -> Result<FoundHosts<&Self>, io::Error> {
lookup_host(&self, qname).await
}
pub async fn search_host(
&self, qname: impl ToRelativeDname
) -> Result<FoundHosts<&Self>, io::Error> {
search_host(&self, qname).await
}
pub async fn lookup_srv(
&self,
service: impl ToRelativeDname,
name: impl ToDname,
fallback_port: u16
) -> Result<Option<FoundSrvs>, SrvError> {
lookup_srv(&self, service, name, fallback_port).await
}
}
#[cfg(feature = "sync")]
impl StubResolver {
pub fn run<R, F>(op: F) -> R::Output
where
R: Future + Send + 'static,
F: FnOnce(StubResolver) -> R + Send + 'static,
{
Self::run_with_conf(ResolvConf::default(), op)
}
pub fn run_with_conf<R, F>(
conf: ResolvConf,
op: F
) -> R::Output
where
R: Future + Send + 'static,
F: FnOnce(StubResolver) -> R + Send + 'static,
{
let resolver = Self::from_conf(conf);
let mut runtime = runtime::Builder::new()
.basic_scheduler()
.build().unwrap();
runtime.block_on(op(resolver))
}
}
impl Default for StubResolver {
fn default() -> Self {
Self::new()
}
}
impl<'a> Resolver for &'a StubResolver {
type Octets = Bytes;
type Answer = Answer;
type Query = Pin<Box<dyn Future<Output = Result<Answer, io::Error>> + 'a>>;
fn query<N, Q>(&self, question: Q) -> Self::Query
where N: ToDname, Q: Into<Question<N>> {
let message = Query::create_message(question.into());
self.query_message(message).boxed()
}
}
impl<'a> SearchNames for &'a StubResolver {
type Name = SearchSuffix;
type Iter = SearchIter<'a>;
fn search_iter(&self) -> Self::Iter {
SearchIter {
resolver: self.clone(),
pos: 0
}
}
}
pub struct Query<'a> {
resolver: &'a StubResolver,
preferred: bool,
attempt: usize,
counter: ServerListCounter,
error: Result<Answer, io::Error>,
}
impl<'a> Query<'a> {
pub fn new(
resolver: &'a StubResolver,
) -> Result<Self, io::Error> {
let (preferred, counter) = if
resolver.options().use_vc ||
resolver.preferred.is_empty()
{
if resolver.stream.is_empty() {
return Err(
io::Error::new(
io::ErrorKind::NotFound,
"no servers available"
)
)
}
(false, resolver.stream.counter(resolver.options().rotate))
}
else {
(true, resolver.preferred.counter(resolver.options().rotate))
};
Ok(Query {
resolver,
preferred,
attempt: 0,
counter,
error: Err(io::Error::new(
io::ErrorKind::TimedOut,
"all timed out"
))
})
}
pub async fn run(
mut self,
mut message: QueryMessage,
) -> Result<Answer, io::Error> {
loop {
match self.run_query(&mut message).await {
Ok(answer) => {
if answer.header().rcode() == Rcode::FormErr
&& self.current_server().does_edns()
{
self.current_server().disable_edns();
continue
}
else if answer.header().rcode() == Rcode::ServFail {
self.update_error_servfail(answer);
}
else if answer.header().tc() && self.preferred
&& !self.resolver.options().ign_tc
{
if self.switch_to_stream() {
continue
}
else {
return Ok(answer)
}
}
else {
return Ok(answer);
}
}
Err(err) => self.update_error(err),
}
if !self.next_server() {
return self.error
}
}
}
fn create_message(
question: Question<impl ToDname>
) -> QueryMessage {
let mut message = MessageBuilder::from_target(
StreamTarget::new(Octets512::new()).unwrap()
).unwrap();
message.header_mut().set_rd(true);
let mut message = message.question();
message.push(question).unwrap();
message.additional()
}
async fn run_query(
&mut self, message: &mut QueryMessage
) -> Result<Answer, io::Error> {
let server = self.current_server();
server.prepare_message(message);
server.query(message).await
}
fn current_server(&self) -> &ServerInfo {
let list = if self.preferred { &self.resolver.preferred }
else { &self.resolver.stream };
self.counter.info(list)
}
fn update_error(&mut self, err: io::Error) {
if err.kind() != io::ErrorKind::TimedOut && self.error.is_err() {
self.error = Err(err)
}
}
fn update_error_servfail(&mut self, answer: Answer) {
self.error = Ok(answer)
}
fn switch_to_stream(&mut self) -> bool {
if !self.preferred {
return false
}
self.preferred = false;
self.attempt = 0;
self.counter = self.resolver.stream.counter(
self.resolver.options().rotate
);
true
}
fn next_server(&mut self) -> bool {
if self.counter.next() {
return true
}
self.attempt += 1;
if self.attempt >= self.resolver.options().attempts {
return false
}
self.counter = if self.preferred {
self.resolver.preferred.counter(self.resolver.options().rotate)
}
else {
self.resolver.stream.counter(self.resolver.options().rotate)
};
true
}
}
pub(super) type QueryMessage = AdditionalBuilder<StreamTarget<Octets512>>;
#[derive(Clone)]
pub struct Answer {
message: Message<Bytes>,
}
impl Answer {
pub fn is_final(&self) -> bool {
(self.message.header().rcode() == Rcode::NoError
|| self.message.header().rcode() == Rcode::NXDomain)
&& !self.message.header().tc()
}
pub fn is_truncated(&self) -> bool {
self.message.header().tc()
}
pub fn into_message(self) -> Message<Bytes> {
self.message
}
}
impl From<Message<Bytes>> for Answer {
fn from(message: Message<Bytes>) -> Self {
Answer { message }
}
}
impl ops::Deref for Answer {
type Target = Message<Bytes>;
fn deref(&self) -> &Self::Target {
&self.message
}
}
impl AsRef<Message<Bytes>> for Answer {
fn as_ref(&self) -> &Message<Bytes> {
&self.message
}
}
#[derive(Clone, Debug)]
pub struct SearchIter<'a> {
resolver: &'a StubResolver,
pos: usize,
}
impl<'a> Iterator for SearchIter<'a> {
type Item = SearchSuffix;
fn next(&mut self) -> Option<Self::Item> {
if let Some(res) = self.resolver.options().search.get(self.pos) {
self.pos += 1;
Some(res.clone())
}
else {
None
}
}
}