diff --git a/Cargo.toml b/Cargo.toml index 289c20cf..a5f2b6d9 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,2 +1,2 @@ [workspace] -members = ["domain-core"] +members = ["domain-core", "domain-resolv"] diff --git a/domain-resolv/Cargo.toml b/domain-resolv/Cargo.toml index 3f5d47b2..9a51d1d3 100644 --- a/domain-resolv/Cargo.toml +++ b/domain-resolv/Cargo.toml @@ -1,5 +1,5 @@ [package] -name = "domain-resolv-preview" +name = "domain-resolv" version = "0.3.1" authors = ["Martin Hoffmann "] description = "An asynchronous DNS stub resolver." @@ -15,11 +15,12 @@ name = "domain_resolv" path = "src/lib.rs" [dependencies] -futures-preview = "0.3.0-alpha.3" -futures-util-preview = { version="0.3.0-alpha.3", features=["compat", "tokio-compat"] } +futures = "^0.1" +rand = "^0.5" tokio = "^0.1" [dependencies.domain-core] path = "../domain-core" version = "0.3.1" + diff --git a/domain-resolv/README.md b/domain-resolv/README.md index c1095633..427e885c 100644 --- a/domain-resolv/README.md +++ b/domain-resolv/README.md @@ -1,20 +1,7 @@ -# domain-resolv-preview +# domain-resolv An asynchronous DNS stub resolver. -*This crate is currently under development and cannot actually resolve -anything just yet.* - - -## Beware - -This crate uses the upcoming async/await features of the Rust compiler as -well as the revised _futures_ crate currently released as _futures-preview._ -As such, it will only work with a reasonably recent nightly compiler. - -This, of course, will change once async/await and future _futures_ have been -stabilized. - ## Usage @@ -22,7 +9,7 @@ First, add this to your `Cargo.toml`: ```toml [dependencies] -domain-resolv-preview = "0.3" +domain-resolv = "0.3" ``` Then, add this to your crate root: @@ -31,5 +18,3 @@ Then, add this to your crate root: extern crate domain_resolv; ``` -Note that the extern crate name is `domain_resolv`. - diff --git a/domain-resolv/src/conf.rs b/domain-resolv/src/conf.rs index 4a7e1a01..590898ff 100644 --- a/domain-resolv/src/conf.rs +++ b/domain-resolv/src/conf.rs @@ -19,6 +19,7 @@ use std::path::Path; use std::str::{self, FromStr, SplitWhitespace}; use std::time::Duration; use domain_core::bits::name::{self, Dname}; +use super::search::SearchList; //------------ ResolvOptions ------------------------------------------------ @@ -32,7 +33,7 @@ use domain_core::bits::name::{self, Dname}; #[derive(Clone, Debug)] pub struct ResolvOptions { /// Search list for host-name lookup. - pub search: Vec, + pub search: SearchList, /// TODO Sortlist /// sortlist: ?? @@ -131,8 +132,8 @@ pub struct ResolvOptions { /// Use bit-label format for IPv6 reverse lookups. /// - /// This option is only relevant for `lookup_addr()` and is implemented - /// there already. + /// Bit labels have been deprecated and consequently, this option is not + /// implemented. pub use_bstring: bool, /// Use ip6.int instead of the recommended ip6.arpa. @@ -171,7 +172,7 @@ impl Default for ResolvOptions { fn default() -> Self { ResolvOptions { // non-flags: - search: Vec::new(), + search: SearchList::new(), //sortlist, ndots: 1, timeout: Duration::new(5,0), @@ -251,7 +252,7 @@ pub struct ServerConf { /// Size of the message receive buffer in bytes. /// /// This is used for datagram transports only. - pub recv_size: u16, + pub recv_size: usize, } impl ServerConf { @@ -362,7 +363,7 @@ impl ResolvConf { pub fn parse_file>( &mut self, path: P ) -> Result<(), Error> { - let mut file = try!(fs::File::open(path)); + let mut file = fs::File::open(path)?; self.parse(&mut file) } @@ -373,7 +374,7 @@ impl ResolvConf { use std::io::BufRead; for line in io::BufReader::new(reader).lines() { - let line = try!(line); + let line = line?; let line = line.trim_right(); if line.is_empty() || line.starts_with(';') || @@ -384,11 +385,11 @@ impl ResolvConf { let mut words = line.split_whitespace(); let keyword = words.next(); match keyword { - Some("nameserver") => try!(self.parse_nameserver(words)), - Some("domain") => try!(self.parse_domain(words)), - Some("search") => try!(self.parse_search(words)), - Some("sortlist") => try!(self.parse_sortlist(words)), - Some("options") => try!(self.parse_options(words)), + Some("nameserver") => self.parse_nameserver(words)?, + Some("domain") => self.parse_domain(words)?, + Some("search") => self.parse_search(words)?, + Some("sortlist") => self.parse_sortlist(words)?, + Some("options") => self.parse_options(words)?, _ => return Err(Error::ParseError) } } @@ -401,7 +402,7 @@ impl ResolvConf { ) -> Result<(), Error> { use std::net::ToSocketAddrs; - for addr in try!((try!(next_word(&mut words)), 53).to_socket_addrs()) + for addr in (next_word(&mut words)?, 53).to_socket_addrs()? { self.servers.push(ServerConf::new(addr, Transport::Udp)); self.servers.push(ServerConf::new(addr, Transport::Tcp)); @@ -413,15 +414,15 @@ impl ResolvConf { &mut self, mut words: SplitWhitespace ) -> Result<(), Error> { - let domain = try!(Dname::from_str(try!(next_word(&mut words)))); - self.options.search = vec![domain]; + let domain = Dname::from_str(next_word(&mut words)?)?; + self.options.search = domain.into(); no_more_words(words) } fn parse_search(&mut self, words: SplitWhitespace) -> Result<(), Error> { - let mut search = Vec::new(); + let mut search = SearchList::new(); for word in words { - let name = try!(Dname::from_str(word)); + let name = Dname::from_str(word)?; search.push(name) } self.options.search = search; @@ -440,7 +441,7 @@ impl ResolvConf { #[allow(match_same_arms)] fn parse_options(&mut self, words: SplitWhitespace) -> Result<(), Error> { for word in words { - match try!(split_arg(word)) { + match split_arg(word)? { ("debug", None) => { } ("ndots", Some(n)) => { self.options.ndots = n @@ -509,20 +510,20 @@ impl fmt::Display for ResolvConf { fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result { for server in &self.servers { let server = server.addr; - try!("nameserver ".fmt(f)); - if server.port() == 53 { try!(server.ip().fmt(f)); } - else { try!(server.fmt(f)); } - try!("\n".fmt(f)); + f.write_str("nameserver ")?; + if server.port() == 53 { server.ip().fmt(f)?; } + else { server.fmt(f)?; } + "\n".fmt(f)?; } if self.options.search.len() == 1 { - try!(write!(f, "domain {}\n", self.options.search[0])); + write!(f, "domain {}\n", self.options.search[0])?; } else if self.options.search.len() > 1 { - try!("search".fmt(f)); - for name in &self.options.search { - try!(write!(f, " {}", name)); + "search".fmt(f)?; + for name in self.options.search.as_slice() { + write!(f, " {}", name)?; } - try!("\n".fmt(f)); + "\n".fmt(f)?; } // Collect options so we only print them if there are any non-default @@ -568,11 +569,11 @@ impl fmt::Display for ResolvConf { if self.options.no_tld_query { options.push("no-tld-query".into()) } if !options.is_empty() { - try!("options".fmt(f)); + "options".fmt(f)?; for option in options { - try!(write!(f, " {}", option)); + write!(f, " {}", option)?; } - try!("\n".fmt(f)); + "\n".fmt(f)?; } Ok(()) @@ -610,7 +611,7 @@ fn split_arg(s: &str) -> Result<(&str, Option), Error> { match s.find(':') { Some(idx) => { let (left, right) = s.split_at(idx); - Ok((left, Some(try!(usize::from_str_radix(&right[1..], 10))))) + Ok((left, Some(usize::from_str_radix(&right[1..], 10)?))) } None => Ok((s, None)) } diff --git a/domain-resolv/src/lib.rs b/domain-resolv/src/lib.rs index 8c5b80a4..dcf3fc27 100644 --- a/domain-resolv/src/lib.rs +++ b/domain-resolv/src/lib.rs @@ -9,18 +9,15 @@ //! top of [futures] and [tokio]. #![allow(unknown_lints)] // hide clippy-related #allows on stable. -// All the unstable features we need to make this work. -#![feature(arbitrary_self_types, async_await, await_macro, futures_api, pin)] - extern crate domain_core; -extern crate futures; -extern crate futures_util; +#[macro_use] extern crate futures; +extern crate rand; extern crate tokio; -pub use self::conf::ResolvConf; -pub use self::resolver::Resolver; - pub mod conf; +pub mod lookup; pub mod resolver; +pub mod search; + +pub mod net; -mod net; diff --git a/domain-resolv/src/lookup/addr.rs b/domain-resolv/src/lookup/addr.rs new file mode 100644 index 00000000..3b05e950 --- /dev/null +++ b/domain-resolv/src/lookup/addr.rs @@ -0,0 +1,156 @@ +//! Looking up host names for addresses. + +use std::net::{IpAddr, Ipv4Addr, Ipv6Addr}; +use std::str::FromStr; +use futures::{Async, Future, Poll}; +use domain_core::bits::message::RecordIter; +use domain_core::bits::name::{Dname, DnameBuilder, ParsedDname}; +use domain_core::iana::Rtype; +use domain_core::rdata::parsed::Ptr; +use ::conf::ResolvOptions; +use ::resolver::{Answer, Query, QueryError, Resolver}; + + +//------------ lookup_addr --------------------------------------------------- + +/// Creates a future that resolves into the host names for an IP address. +/// +/// The future will query DNS using the resolver represented by `resolv`. +/// It will query DNS only and not consider any other database the system +/// may have. +/// +/// The value returned upon success can be turned into an iterator over +/// host names via its `iter()` method. This is due to lifetime issues. +pub fn lookup_addr(resolv: &Resolver, addr: IpAddr) -> LookupAddr { + let name = dname_from_addr(addr, resolv.options()); + LookupAddr(resolv.query((name, Rtype::Ptr))) +} + + +//------------ LookupAddr ---------------------------------------------------- + +/// The future for [`lookup_addr()`]. +/// +/// [`lookup_addr()`]: fn.lookup_addr.html +pub struct LookupAddr(Query); + +impl Future for LookupAddr { + type Item = FoundAddrs; + type Error = QueryError; + + fn poll(&mut self) -> Poll { + Ok(Async::Ready(FoundAddrs(try_ready!(self.0.poll())))) + } +} + + +//------------ FoundAddrs ---------------------------------------------------- + +/// The success type of the `lookup_addr()` function. +/// +/// The only purpose of this type is to return an iterator over host names +/// via its `iter()` method. +pub struct FoundAddrs(Answer); + +impl FoundAddrs { + /// Returns an iterator over the host names. + pub fn iter(&self) -> FoundAddrsIter { + FoundAddrsIter { + name: self.0.canonical_name(), + answer: self.0.answer().ok().map(|sec| sec.limit_to::()) + } + } +} + +impl IntoIterator for FoundAddrs { + type Item = ParsedDname; + type IntoIter = FoundAddrsIter; + + fn into_iter(self) -> Self::IntoIter { + self.iter() + } +} + +impl<'a> IntoIterator for &'a FoundAddrs { + type Item = ParsedDname; + type IntoIter = FoundAddrsIter; + + fn into_iter(self) -> Self::IntoIter { + self.iter() + } +} + + +//------------ FoundAddrsIter ------------------------------------------------ + +/// An iterator over host names returned by address lookup. +pub struct FoundAddrsIter { + name: Option, + answer: Option>, +} + +impl Iterator for FoundAddrsIter { + type Item = ParsedDname; + + #[allow(while_let_on_iterator)] + fn next(&mut self) -> Option { + let name = if let Some(ref name) = self.name { name } + else { return None }; + let answer = if let Some(ref mut answer) = self.answer { answer } + else { return None }; + while let Some(Ok(record)) = answer.next() { + if record.owner() == name { + return Some(record.into_data().into_ptrdname()) + } + } + None + } +} + + + +//------------ Helper Functions --------------------------------------------- + +/// Translates an IP address into a domain name. +fn dname_from_addr(addr: IpAddr, opts: &ResolvOptions) -> Dname { + match addr { + IpAddr::V4(addr) => dname_from_v4(addr), + IpAddr::V6(addr) => dname_from_v6(addr, opts) + } +} + +/// Translates an IPv4 address into a domain name. +fn dname_from_v4(addr: Ipv4Addr) -> Dname { + // XXX There’s a more efficient way to doing this. + let octets = addr.octets(); + Dname::from_str( + &format!( + "{}.{}.{}.{}.in-addr.arpa.", octets[3], + octets[2], octets[1], octets[0]) + ).unwrap() +} + +/// Translate an IPv6 address into a domain name. +/// +/// As there are several ways to do this, the functions depends on +/// resolver options, namely `use_bstring` and `use_ip6dotin`. +fn dname_from_v6(addr: Ipv6Addr, opts: &ResolvOptions) -> Dname { + let mut res = DnameBuilder::new(); + for item in addr.segments().iter().rev() { + let text = format!("{:04x}", item); + let text = text.as_bytes(); + res.append_label(&text[3..4]).unwrap(); + res.append_label(&text[2..3]).unwrap(); + res.append_label(&text[1..2]).unwrap(); + res.append_label(&text[0..1]).unwrap(); + } + res.append_label(b"ip6").unwrap(); + if opts.use_ip6dotint { + res.append_label(b"int").unwrap(); + } + else { + res.append_label(b"arpa").unwrap(); + } + res.into_dname().unwrap() +} + diff --git a/domain-resolv/src/lookup/host.rs b/domain-resolv/src/lookup/host.rs new file mode 100644 index 00000000..91fc637f --- /dev/null +++ b/domain-resolv/src/lookup/host.rs @@ -0,0 +1,275 @@ +//! Looking up host names. + +use std::{io, mem, slice}; +use std::net::{IpAddr, SocketAddr, ToSocketAddrs}; +use domain_core::bits::name::{ + Dname, ParsedDname, ParsedDnameError, ToDname, ToRelativeDname +}; +use domain_core::iana::Rtype; +use domain_core::rdata::parsed::{A, Aaaa}; +use tokio::prelude::{Async, Future, Poll}; +use ::resolver::{Answer, Query, QueryError, Resolver}; +use ::search::search; + + +//------------ lookup_host --------------------------------------------------- + +/// Creates a future that resolves a host name into its IP addresses. +/// +/// The future will use the resolver given in `resolv` to query the +/// DNS for the IPv4 and IPv6 addresses associated with `name`. If `name` +/// is a relative domain name, it is being translated into a series of +/// absolute names according to the resolver’s configuration. +/// +/// The value returned upon success can be turned into an iterator over +/// IP addresses or even socket addresses. Since the lookup may determine that +/// the host name is in fact an alias for another name, the value will also +/// return the canonical name. +pub fn lookup_host(resolver: &Resolver, name: &N) -> LookupHost { + LookupHost { + a: MaybeDone::NotYet(resolver.query((name, Rtype::A))), + aaaa: MaybeDone::NotYet(resolver.query((name, Rtype::Aaaa))), + } +} + +pub fn search_host( + resolver: &Resolver, + name: N +) -> impl Future { + search(resolver, name, |resolver, name| lookup_host(resolver, &name)) +} + + +//------------ LookupHost ---------------------------------------------------- + +/// The future for [`lookup_host()`]. +/// +/// [`lookup_host()`]: fn.lookup_host.html +#[derive(Debug)] +pub struct LookupHost { + /// The A query for the currently processed name. + a: MaybeDone, + + /// The AAAA query for the currently processed name. + aaaa: MaybeDone, +} + + +//--- Future + +impl Future for LookupHost { + type Item = FoundHosts; + type Error = QueryError; + + fn poll(&mut self) -> Poll { + if (self.a.poll(), self.aaaa.poll()) != (true, true) { + return Ok(Async::NotReady) + } + match FoundHosts::from_answers(self.a.take(), self.aaaa.take()) { + Ok(res) => Ok(Async::Ready(res)), + Err(err) => Err(err) + } + } +} + + +//------------ MaybeDone ----------------------------------------------------- + +/// A future that may or may not yet have been resolved. +/// +/// This is mostly the type used by futures’ own `join()`, except that we need +/// to consider errors as part success is still good. +#[derive(Debug)] +enum MaybeDone { + /// It’s still ongoing. + NotYet(A), + + /// It resolved successfully. + Item(A::Item), + + /// It resolved with an error. + Error(A::Error), + + /// It is gone. + Gone +} + +impl MaybeDone { + /// Polls the wrapped future. + /// + /// Returns whether the future is resolved. + /// + /// # Panics + /// + /// If the value is of the `MaybeDone::Gone`, calling this function will + /// panic. + fn poll(&mut self) -> bool { + let res = match *self { + MaybeDone::NotYet(ref mut a) => a.poll(), + MaybeDone::Item(_) | MaybeDone::Error(_) => return true, + MaybeDone::Gone => panic!("polling a completed LookupHost"), + }; + match res { + Ok(Async::Ready(item)) => { + *self = MaybeDone::Item(item); + true + } + Err(err) => { + *self = MaybeDone::Error(err); + true + } + Ok(Async::NotReady) => { + false + } + } + } + + /// Trades the value in for the result of the future. + /// + /// # Panics + /// + /// Panics if there isn’t actually a result. + fn take(&mut self) -> Result { + match mem::replace(self, MaybeDone::Gone) { + MaybeDone::Item(item) => Ok(item), + MaybeDone::Error(err) => Err(err), + _ => panic!(), + } + } +} + + + + +//------------ FoundHosts ---------------------------------------------------- + +/// The value returned by a successful host lookup. +/// +/// You can use the `iter()` method to get an iterator over the IP addresses +/// or `port_iter()` to get an iterator over socket addresses with the given +/// port. +/// +/// The `canonical_name()` method returns the canonical name of the host for +/// which the addresses were found. +#[derive(Clone, Debug)] +pub struct FoundHosts { + /// The canonical domain name for the host. + canonical: Dname, + + /// All the IP addresses we’ve got. + addrs: Vec +} + +impl FoundHosts { + pub fn new(canonical: Dname, addrs: Vec) -> Self { + FoundHosts { canonical: canonical, addrs: addrs } + } + + /// Creates a new value from the results of the A and AAAA queries. + /// + /// Either of the queries can have resulted in an error but not both. + fn from_answers( + a: Result, b: Result + ) -> Result { + let (a, b) = match (a, b) { + (Ok(a), b) => (a, b), + (a, Ok(b)) => (b, a), + (Err(a), Err(b)) => return Err(a.merge(b)) + }; + let name = a.canonical_name().unwrap(); + let mut addrs = Vec::new(); + Self::process_records(&mut addrs, &a, &name).ok(); + if let Ok(b) = b { + Self::process_records(&mut addrs, &b, &name).ok(); + } + Ok(FoundHosts { + canonical: name.to_name(), + addrs: addrs + }) + } + + /// Processes the records of a response message. + /// + /// Adds all A and AAA records contained in `msg`’s answer to `addrs`, + /// assuming they domain name in the record matches `name`. + fn process_records( + addrs: &mut Vec, + msg: &Answer, + name: &ParsedDname + ) -> Result<(), ParsedDnameError> { + for record in msg.answer()?.limit_to::() { + if let Ok(record) = record { + if record.owner() == name { + addrs.push(IpAddr::V4(record.data().addr())) + } + } + } + for record in msg.answer()?.limit_to::() { + if let Ok(record) = record { + if record.owner() == name { + addrs.push(IpAddr::V6(record.data().addr())) + } + } + } + Ok(()) + } + + /// Returns a reference to the canonical name for the host. + pub fn canonical_name(&self) -> &Dname { + &self.canonical + } + + /// Returns an iterator over the IP addresses returned by the lookup. + pub fn iter(&self) -> FoundHostsIter { + FoundHostsIter(self.addrs.iter()) + } + + /// Returns an iterator over socket addresses gained from the lookup. + /// + /// The socket addresses are gained by combining the IP addresses with + /// `port`. The returned iterator implements `ToSocketAddrs` and thus + /// can be used where `std::net` wants addresses right away. + pub fn port_iter(&self, port: u16) -> FoundHostsSocketIter { + FoundHostsSocketIter(self.addrs.iter(), port) + } +} + + +//------------ FoundHostsIter ------------------------------------------------ + +/// An iterator over the IP addresses returned by a host lookup. +#[derive(Clone, Debug)] +pub struct FoundHostsIter<'a>(slice::Iter<'a, IpAddr>); + +impl<'a> Iterator for FoundHostsIter<'a> { + type Item = IpAddr; + + fn next(&mut self) -> Option { + self.0.next().cloned() + } +} + + +//------------ FoundHostsSocketIter ------------------------------------------ + +/// An iterator over socket addresses derived from a host lookup. +#[derive(Clone, Debug)] +pub struct FoundHostsSocketIter<'a>(slice::Iter<'a, IpAddr>, u16); + +impl<'a> Iterator for FoundHostsSocketIter<'a> { + type Item = SocketAddr; + + fn next(&mut self) -> Option { + self.0.next().map(|addr| SocketAddr::new(*addr, self.1)) + } +} + +impl<'a> ToSocketAddrs for FoundHostsSocketIter<'a> { + type Iter = Self; + + fn to_socket_addrs(&self) -> io::Result { + Ok(self.clone()) + } +} + + diff --git a/domain-resolv/src/lookup/mod.rs b/domain-resolv/src/lookup/mod.rs new file mode 100644 index 00000000..00406375 --- /dev/null +++ b/domain-resolv/src/lookup/mod.rs @@ -0,0 +1,13 @@ +//! Lookup functions and related types. +//! +//! This module collects a number of more or less complex lookups that +//! implement applications of the DNS. + +pub use self::addr::lookup_addr; +pub use self::host::lookup_host; +pub use self::srv::lookup_srv; + +pub mod addr; +pub mod host; +pub mod srv; + diff --git a/domain-resolv/src/lookup/records.rs b/domain-resolv/src/lookup/records.rs new file mode 100644 index 00000000..d2448b72 --- /dev/null +++ b/domain-resolv/src/lookup/records.rs @@ -0,0 +1,28 @@ +//! Looking up raw records. + +use domain_core::bits::name::ToDname; +use domain_core::bits::Question; +use tokio::prelude::Future; +use crate::resolver::{Answer, QueryError, Resolver}; + + +//------------ lookup_records ------------------------------------------------ + +/// Creates a future that looks up DNS records. +/// +/// The future will use the given resolver to perform a DNS query for the +/// records of type `rtype` associated with `name` in `class`. +/// This differs from calling `resolv.query()` directly in that it can treat +/// relative names. In this case, the resolver configuration is considered +/// to translate the name into a series of absolute names. If you want to +/// find out the name that resulted in a successful answer, you can look at +/// the query in the resulting message. +pub fn lookup_records<'a, N, Q>( + resolver: &'a Resolver, + question: Q +) -> impl Future> + 'a +where N: ToDname + 'a, Q: Into>+ 'a { + resolver.query(question) +} + + diff --git a/domain-resolv/src/lookup/srv.rs b/domain-resolv/src/lookup/srv.rs new file mode 100644 index 00000000..1b6dfb19 --- /dev/null +++ b/domain-resolv/src/lookup/srv.rs @@ -0,0 +1,491 @@ +//! Looking up SRV records. + +use domain_core::bits::name::{ + Dname, ParsedDname, ParsedDnameError, ToRelativeDname, ToDname +}; +use domain_core::iana::Rtype; +use domain_core::rdata::parsed::{A, Aaaa, Srv}; +use rand; +use rand::distributions::{Distribution, Range}; +use tokio::prelude::{Async, Future, Poll, Stream}; +use ::resolver::{Answer, Query, QueryError, Resolver}; +use super::host::{FoundHosts, FoundHostsSocketIter, LookupHost, lookup_host}; + + +//------------ lookup_records ------------------------------------------------ + +/// Creates a future that looks up SRV records. +/// +/// The future will use the resolver given in `resolver` to query the +/// DNS for SRV records associated with domain name `name` and service +/// `service`. +/// +/// The value returned upon success can be turned into a stream of +/// `ResolvedSrvItem`s corresponding to the found SRV records, ordered as per +/// the usage rules defined in [RFC 2782]. If no matching SRV record is found, +/// A/AAAA queries on the bare domain name `name` will be attempted, yielding +/// a single element upon success using the port given by `fallback_port`, +/// typcially the standard port for the service in question. +/// +/// Each item in the stream can be turned into an iterator over socket +/// addresses as accepted by, for instance, `TcpStream::connect`. +/// +/// The future resolves to `None` whenever the request service is +/// “decidedly not available” at the requested domain, that is there is a +/// single SRV record with the root label as its target. +pub fn lookup_srv( + resolver: &Resolver, + service: S, + name: N, + fallback_port: u16 +) -> LookupSrv +where + S: ToRelativeDname + Clone + Send + 'static, + N: ToDname + Send + 'static +{ + let query = { + let full_name = match (&service).chain(&name) { + Ok(name) => name, + Err(_) => { + return LookupSrv { + data: None, + query: Err(Some(SrvError::LongName)) + } + } + }; + resolver.query((full_name, Rtype::Srv)) + }; + LookupSrv { + data: Some(LookupData { + resolver: resolver.clone(), + host: name, + service, + fallback_port + }), + query: Ok(query) + } +} + + +//------------ LookupData ---------------------------------------------------- + +#[derive(Debug)] +struct LookupData { + /// The resolver to run queries on. + resolver: Resolver, + + /// Bare host to be queried. + /// + /// This is kept for fallback if no SRV records are found. + host: N, + + /// Service name + service: S, + + /// Fallback port, used if no SRV records are found + fallback_port: u16, +} + + +//------------ LookupSrv ----------------------------------------------------- + +/// The future returned by [`lookup_srv()`]. +/// +/// [`lookup_srv()`]: fn.lookup_srv.html +pub struct LookupSrv { + data: Option>, + query: Result>, +} + + +impl Future for LookupSrv +where + S: ToRelativeDname + Clone + Send + 'static, + N: ToDname + Send + 'static +{ + type Item = Option>; + type Error = SrvError; + + fn poll(&mut self) -> Poll { + match self.query { + Ok(ref mut query) => match query.poll() { + Ok(Async::NotReady) => Ok(Async::NotReady), + Ok(Async::Ready(answer)) => { + Ok(Async::Ready( + FoundSrvs::new( + answer, + self.data.take().expect("polled resolved future") + )? + )) + } + Err(_) => { + Ok(Async::Ready(Some( + FoundSrvs::new_dummy( + self.data.take().expect("polled resolved future")) + ))) + } + } + Err(ref mut err) => { + Err(err.take().expect("polled resolved future")) + } + } + } +} + + +//------------ LookupSrvStream ----------------------------------------------- + +/// Stream over SrvItem elements. +/// +/// SrvItem elements are resolved as needed, skipping them in case of failure. +/// It is therefore guaranteed to yield only SrvItem structs that have +/// a `SrvItemState::Resolved` state. +#[derive(Debug)] +pub struct LookupSrvStream { + /// The resolver to use for A/AAAA requests. + resolver: Resolver, + + /// A vector of (potentially unresolved) SrvItem elements. + /// + /// Note that we take items from this via `pop`, so it needs to be ordered + /// backwards. + items: Vec>, + + /// A/AAAA lookup for the last `SrvItem` in `items`. + lookup: Option +} + +impl LookupSrvStream { + fn new(found: FoundSrvs) -> Self { + LookupSrvStream { + resolver: found.resolver, + items: found.items.into_iter().rev().collect(), + lookup: None, + } + } +} + + +//--- Stream + +impl Stream for LookupSrvStream +where S: ToRelativeDname + Clone + Send + 'static { + type Item = ResolvedSrvItem; + type Error = SrvError; + + fn poll(&mut self) -> Poll, Self::Error> { + // See if we have a query result. We need to break this in to because + // of the mut ref on the inside of self.lookup. + let res = if let Some(ref mut query) = self.lookup { + match query.poll() { + Ok(Async::NotReady) => return Ok(Async::NotReady), + Ok(Async::Ready(found)) => { + Some(ResolvedSrvItem::from_item_and_hosts( + self.items.pop().unwrap(), + found + )) + } + Err(_) => None + } + } + else { + None + }; + + // We have a query result. Clear lookup and return. + if let Some(res) = res { + self.lookup = None; + return Ok(Async::Ready(Some(res))) + } + + // Start a new query if necessary. Return if we are done. + match self.items.last() { + Some(item) => match item.state { + SrvItemState::Unresolved(ref host) => { + self.lookup = Some(lookup_host(&self.resolver, host)); + } + _ => { } + } + None => return Ok(Async::Ready(None)) // we are done. + } + + if self.lookup.is_some() { + self.poll() + } + else { + Ok(Async::Ready(Some( + ResolvedSrvItem::from_item(self.items.pop().unwrap()).unwrap() + ))) + } + } +} + + +//------------ FoundSrvs ----------------------------------------------------- + +#[derive(Clone, Debug)] +pub struct FoundSrvs { + resolver: Resolver, + items: Vec>, +} + +impl FoundSrvs { + pub fn into_stream(self) -> LookupSrvStream { + LookupSrvStream::new(self) + } + + /// Moves all results from `other` into `Self`, leaving `other` empty. + /// + /// Reorders merged results as if they were from a single query. + pub fn merge(&mut self, other : &mut Self) { + self.items.append(&mut other.items); + Self::reorder_items(&mut self.items); + } +} + +impl FoundSrvs { + fn new( + answer: Answer, + data: LookupData + ) -> Result, SrvError> { + let name = answer.canonical_name().unwrap(); + let mut rrs = Vec::new(); + Self::process_records(&mut rrs, &answer, &name)?; + + if rrs.len() == 0 { + return Ok(Some(Self::new_dummy(data))) + } + if rrs.len() == 1 && rrs[0].target().is_root() { + // Exactly one record with target "." indicates no service. + return Ok(None) + } + + // Build results including potentially resolved IP addresses + let mut items = Vec::with_capacity(rrs.len()); + Self::items_from_rrs(&rrs, &answer, &mut items, &data)?; + Self::reorder_items(&mut items); + + Ok(Some(FoundSrvs { + resolver: data.resolver, + items + })) + } + + fn new_dummy(data: LookupData) -> Self { + FoundSrvs { + resolver: data.resolver, + items: vec![ + SrvItem { + priority: 0, + weight: 0, + port: data.fallback_port, + service: None, + state: SrvItemState::Unresolved(data.host.to_name()) + } + ] + } + } + + fn process_records( + rrs: &mut Vec, + answer: &Answer, + name: &ParsedDname + ) -> Result<(), SrvError> { + for record in answer.answer()?.limit_to::() { + if let Ok(record) = record { + if record.owner() == name { + rrs.push(record.data().clone()) + } + } + } + Ok(()) + } + + fn items_from_rrs( + rrs: &[Srv], + answer: &Answer, + result: &mut Vec>, + data: &LookupData, + ) -> Result<(), SrvError> { + for rr in rrs { + let mut addrs = Vec::new(); + let name = rr.target().to_name(); + for record in answer.additional()?.limit_to::() { + if let Ok(record) = record { + if record.owner() == &name { + addrs.push(record.data().addr().into()) + } + } + } + for record in answer.additional()?.limit_to::() { + if let Ok(record) = record { + if record.owner() == &name { + addrs.push(record.data().addr().into()) + } + } + } + let state = if addrs.is_empty() { + SrvItemState::Unresolved(name) + } + else { + SrvItemState::Resolved(FoundHosts::new(name, addrs)) + }; + result.push(SrvItem { + priority: rr.priority(), + weight: rr.weight(), + state: state, + port: rr.port(), + service: Some(data.service.clone()) + }) + } + Ok(()) + } +} + +impl FoundSrvs { + fn reorder_items(items: &mut [SrvItem]) { + // First, reorder by priority and weight, effectively + // grouping by priority, with weight 0 records at the beginning of + // each group. + items.sort_by_key(|k| (k.priority, k.weight)); + + // Find each group and reorder them using reorder_by_weight + let mut current_prio = 0; + let mut weight_sum = 0; + let mut first_index = 0; + for i in 0 .. items.len() { + if current_prio != items[i].priority { + current_prio = items[i].priority; + Self::reorder_by_weight(&mut items[first_index..i], weight_sum); + weight_sum = 0; + first_index = i; + } + weight_sum += items[i].weight as u32; + } + Self::reorder_by_weight(&mut items[first_index..], weight_sum); + } + + /// Reorders items in a priority level based on their weight + fn reorder_by_weight(items: &mut [SrvItem], weight_sum : u32) { + let mut rng = rand::thread_rng(); + let mut weight_sum = weight_sum; + for i in 0 .. items.len() { + let range = Range::new(0, weight_sum + 1); + let mut sum : u32 = 0; + let pick = range.sample(&mut rng); + for j in 0 .. items.len() { + sum += items[j].weight as u32; + if sum >= pick { + weight_sum -= items[j].weight as u32; + items.swap(i, j); + break; + } + } + } + } +} + + +//------------ SrvItem ------------------------------------------------------- + +#[derive(Clone, Debug)] +pub struct SrvItem { + priority: u16, + weight: u16, + port: u16, + service: Option, + state: SrvItemState +} + +#[derive(Clone, Debug)] +pub enum SrvItemState { + Unresolved(Dname), + Resolved(FoundHosts) +} + +impl SrvItem { + + /// Returns a reference to the service + proto part of the domain name. + /// + /// Useful when mixing results from different SRV queries. + pub fn txt_service(&self) -> Option<&S> { + self.service.as_ref() + } + + /// Returns a reference to the name of the target. + pub fn target(&self) -> &Dname { + match self.state { + SrvItemState::Unresolved(ref target) => target, + SrvItemState::Resolved(ref found_hosts) => found_hosts.canonical_name() + } + } +} + + +//------------ ResolvedSrvItem ----------------------------------------------- + +#[derive(Clone, Debug)] +pub struct ResolvedSrvItem { + priority: u16, + weight: u16, + port: u16, + service: Option, + hosts: FoundHosts, +} + +impl ResolvedSrvItem { + /// Returns an iterator over socket addresses matching an SRV record. + /// + /// SrvItem does not implement the `ToSocketAddrs` trait as the result + /// of `to_socket_addrs()` does not have a static lifetime. + pub fn to_socket_addrs(&self) -> FoundHostsSocketIter { + self.hosts.port_iter(self.port) + } + + fn from_item(item: SrvItem) -> Option { + if let SrvItemState::Resolved(hosts) = item.state { + Some(ResolvedSrvItem { + priority: item.priority, + weight: item.weight, + port: item.port, + service: item.service, + hosts: hosts + }) + } + else { + None + } + } + + fn from_item_and_hosts(item: SrvItem, hosts: FoundHosts) -> Self { + ResolvedSrvItem { + priority: item.priority, + weight: item.weight, + port: item.port, + service: item.service, + hosts: hosts + } + } +} + + +//------------ SrvError ------------------------------------------------------ + +#[derive(Debug)] +pub enum SrvError { + LongName, + Query(QueryError), +} + +impl From for SrvError { + fn from(err: QueryError) -> SrvError { + SrvError::Query(err) + } +} + +impl From for SrvError { + fn from(_: ParsedDnameError) -> SrvError { + SrvError::Query(QueryError::MalformedAnswer) + } +} + diff --git a/domain-resolv/src/net.rs b/domain-resolv/src/net.rs deleted file mode 100644 index b2942d23..00000000 --- a/domain-resolv/src/net.rs +++ /dev/null @@ -1,256 +0,0 @@ -//! Networking. -//! -//! This private module takes care of all the asynchronous networking. It is a -//! bit messy currently due to having to deal with compatibility between -//! futures 0.1 and 0.3 for tokio. - -use std::io; -use std::net::SocketAddr; -use std::sync::atomic::{AtomicUsize, Ordering}; -use std::time::Instant; -use domain_core::bits::Message; -use domain_core::bits::message_builder::OptBuilder; -use futures::future::TryFutureExt; -use futures_util::compat::Future01CompatExt; -use tokio::io::{read_exact, write_all}; -use tokio::net::{TcpStream, UdpSocket}; -use tokio::timer::Delay; -use super::conf::{ResolvConf, ServerConf, Transport}; -use super::resolver::Answer; - - -//------------ Module Configuration ------------------------------------------ - -/// How many times do we try a new random port if we get ‘address in use.’ -const RETRY_RANDOM_PORT: usize = 10; - - -//------------ Async Networking Functions ------------------------------------ - -pub async fn query_server( - server: &ServerConf, - mut message: OptBuilder -) -> (OptBuilder, Result) { - message.set_udp_payload_size(server.recv_size); - message.header_mut().set_random_id(); - let res = { - let fut = _query_server(server.transport, server.addr, &message); - await!(fut.try_join(delay(Instant::now() + server.request_timeout))) - }; - (message, res.unwrap_err()) -} - -/// The future for the timeout. -/// -/// The weird return type is to trick `TryFutureExt::try_join` into returning -/// as soon as the first future completes. -async fn delay(instant: Instant) -> Result<(), Result> { - // Delay completes with Ok(()) when the timeout is reached or some error - // when things go awry. We deal with the latter by also just timing out. - // So, however Delay completes, a timeout error is the response. - await!(Delay::new(instant).compat()); - Err(Err( - io::Error::new( - io::ErrorKind::TimedOut, - "timed out" - ) - )) -} - -/// The future for actually querying a server. -/// -/// Separatedly takes the transport and (remote) address instead of the -/// complete server config to avoid having two references as arguments which -/// would require lifetime parameters which currently crashes rustc. -/// -/// The weird return type is to trick `TryFutureExt::try_join` into returning -/// as soon as the first future completes. -async fn _query_server( - transport: Transport, - addr: SocketAddr, - message: &OptBuilder -) -> Result<(), Result> { - let res = match transport { - Transport::Udp => await!(query_udp_server(addr, &message)), - Transport::Tcp => await!(query_tcp_server(addr, &message)), - }; - Err(res) -} - -/// The future for querying a UDP server. -async fn query_udp_server( - addr: SocketAddr, - message: &OptBuilder -) -> Result { - let sock = bind_udp(addr.is_ipv4())?; - println!("got socket"); - sock.connect(&addr)?; - println!("connected socket"); - let (mut sock, _) = await!( - sock.send_dgram(&message.preview()[2..], &addr) .compat() - )?; - println!("sent message"); - loop { - let buf = vec![0; 4096]; // XXX Or what? - let (the_sock, mut buf, size, _addr) = await!( - sock.recv_dgram(buf).compat() - )?; - println!("received message"); - sock = the_sock; - buf.truncate(size); - let answer = match Message::from_bytes(buf.into()) { - Ok(msg) => msg, - Err(_) => { - return Err(io::Error::new(io::ErrorKind::Other, "short buf")) - } - }; - println!("parsed message"); - // XXX Check question. - if is_answer(message, &answer) { - return Ok(answer.into()) - } - println!("not an answer"); - } -} - -/// Creates a bound UDP socket. -/// -/// We are supposed to pick a random local port for socket for extra -/// protection. So we try just that here. -fn bind_udp(v4: bool) -> Result { - let mut i = 0; - loop { - let local = if v4 { ([0u8; 4], 0).into() } - else { ([0u16; 8], 0).into() }; - match UdpSocket::bind(&local) { - Ok(sock) => return Ok(sock), - Err(err) => { - if i == RETRY_RANDOM_PORT { - return Err(err); - } - else { - i += 1 - } - } - } - } -} - -/// The future that queries a TCP server. -async fn query_tcp_server( - addr: SocketAddr, - message: &OptBuilder -) -> Result { - let sock = await!(TcpStream::connect(&addr).compat())?; - let (mut sock, _) = await!(write_all(sock, message.preview()).compat())?; - loop { - let (res_sock, buf) = await!(read_exact(sock, [0u8; 2]).compat())?; - let len = (buf[0] as usize) << 8 | buf[1] as usize; - let (res_sock, buf) = await!( - read_exact(res_sock, vec![0u8; len]).compat() - )?; - let answer = Message::from_bytes(buf.into()).map_err(|_| { - io::Error::new(io::ErrorKind::Other, "short buf") - })?; - if is_answer(message, &answer) { - return Ok(answer.into()) - } - sock = res_sock; - } -} - - -fn is_answer(request: &OptBuilder, response: &Message) -> bool { - // XXX Also compare the questions, but we need to figure out how to do - // that for those two types. - println!("{} {}", request.header().id(), response.header().id()); - request.header().id() == response.header().id() -} - - -//------------ ServerList ---------------------------------------------------- - -#[derive(Debug)] -pub struct ServerList { - /// The actual list of servers. - servers: Vec, - - /// Where to start accessing the list. - /// - /// This value will always keep growing and will have to be used module - /// `servers`’s length. - /// - /// When it eventually wraps around the end of usize’s range, there will - /// be a jump in rotation. Since that will happen only oh-so-often, we - /// accept that in favour of simpler code. - start: AtomicUsize, -} - -impl ServerList { - pub fn from_conf(conf: &ResolvConf, filter: F) -> Self - where F: Fn(&ServerConf) -> bool { - ServerList { - servers: { - conf.servers.iter().filter(|f| filter(*f)) - .map(Clone::clone).collect() - }, - start: AtomicUsize::new(0), - } - } - - pub fn iter(&self) -> ServerListIter { - ServerListIter::new(self) - } - - pub fn rotate(&self) { - self.start.fetch_add(1, Ordering::SeqCst); - } -} - -impl<'a> IntoIterator for &'a ServerList { - type Item = &'a ServerConf; - type IntoIter = ServerListIter<'a>; - - fn into_iter(self) -> Self::IntoIter { - self.iter() - } -} - - -//------------ ServerListIter ------------------------------------------------ - -#[derive(Clone, Debug)] -pub struct ServerListIter<'a> { - servers: &'a [ServerConf], - cur: usize, - end: usize, -} - -impl<'a> ServerListIter<'a> { - fn new(list: &'a ServerList) -> Self { - // We modulo the start value here to prevent hick-ups towards the - // end of usize’s range. - let start = list.start.load(Ordering::Relaxed) % list.servers.len(); - ServerListIter { - servers: list.servers.as_ref(), - cur: start, - end: start + list.servers.len(), - } - } -} - -impl<'a> Iterator for ServerListIter<'a> { - type Item = &'a ServerConf; - - fn next(&mut self) -> Option { - if self.cur == self.end { - None - } - else { - let res = &self.servers[self.cur % self.servers.len()]; - self.cur += 1; - Some(res) - } - } -} - diff --git a/domain-resolv/src/net/mod.rs b/domain-resolv/src/net/mod.rs new file mode 100644 index 00000000..e825596d --- /dev/null +++ b/domain-resolv/src/net/mod.rs @@ -0,0 +1,250 @@ +use std::{io, ops}; +use std::net::SocketAddr; +use std::sync::Arc; +use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; +use domain_core::bits::query::{QueryBuilder, QueryMessage}; +use tokio::prelude::{Async, Future}; +use tokio::timer::Timeout; +use super::conf::{ResolvConf, ServerConf, Transport}; +use super::resolver::Answer; + +mod tcp; +mod udp; +mod util; + + +//------------ ServerInfo ---------------------------------------------------- + +#[derive(Debug)] +pub struct ServerInfo { + /// The basic server configuration. + conf: ServerConf, + + /// Whether this server supports EDNS. + /// + /// We start out with assuming it does and unset it if we get a FORMERR. + edns: AtomicBool, +} + +impl ServerInfo { + pub fn conf(&self) -> &ServerConf { + &self.conf + } + + pub fn does_edns(&self) -> bool { + self.edns.load(Ordering::Relaxed) + } + + pub fn disable_edns(&self) { + self.edns.store(false, Ordering::Relaxed); + } + + pub fn prepare_message(&self, query: &mut QueryBuilder) { + query.revert_additional(); + if self.does_edns() { + query.add_opt(|opt| { + // These are the values that Unbound uses. + // XXX Perhaps this should be configurable. + opt.header_mut().set_udp_payload_size( + match self.conf.addr { + SocketAddr::V4(_) => 1472, + SocketAddr::V6(_) => 1232 + } + ) + }) + } + } +} + +impl From for ServerInfo { + fn from(conf: ServerConf) -> Self { + ServerInfo { + conf, + edns: AtomicBool::new(true) + } + } +} + +impl<'a> From<&'a ServerConf> for ServerInfo { + fn from(conf: &'a ServerConf) -> Self { + conf.clone().into() + } +} + + +//------------ ServerQuery --------------------------------------------------- + +#[derive(Debug)] +pub enum ServerQuery { + Tcp(Timeout), + Udp(Timeout), +} + +impl ServerQuery { + pub fn new(query: QueryMessage, server: &ServerInfo) -> Self { + match server.conf.transport { + Transport::Udp => { + ServerQuery::Udp(Timeout::new( + udp::UdpQuery::new( + query, + server.conf.addr, + server.conf.recv_size, + ), + server.conf.request_timeout + )) + } + Transport::Tcp => { + ServerQuery::Tcp(Timeout::new( + tcp::TcpQuery::new(query, server.conf.addr), + server.conf.request_timeout + )) + } + } + } +} + +impl Future for ServerQuery { + type Item = Answer; + type Error = io::Error; + + fn poll(&mut self) -> Result, Self::Error> { + match *self { + ServerQuery::Tcp(ref mut tcp) => tcp.poll(), + ServerQuery::Udp(ref mut udp) => udp.poll(), + }.map_err(|err| { + err.into_inner().unwrap_or_else(|| + io::Error::new(io::ErrorKind::TimedOut, "timed out") + ) + }) + } +} + + +//------------ ServerList ---------------------------------------------------- + +#[derive(Debug)] +pub struct ServerList { + /// The actual list of servers. + servers: Vec, + + /// Where to start accessing the list. + /// + /// In rotate mode, this value will always keep growing and will have to + /// be used modulo `servers`’s length. + /// + /// When it eventually wraps around the end of usize’s range, there will + /// be a jump in rotation. Since that will happen only oh-so-often, we + /// accept that in favour of simpler code. + start: Arc, +} + +impl ServerList { + pub fn from_conf(conf: &ResolvConf, filter: F) -> Self + where F: Fn(&ServerConf) -> bool { + ServerList { + servers: { + conf.servers.iter().filter(|f| filter(*f)) + .map(Into::into).collect() + }, + start: Arc::new(AtomicUsize::new(0)), + } + } + + pub fn counter(&self, rotate: bool) -> ServerListCounter { + let res = ServerListCounter::new(self); + if rotate { + self.rotate() + } + res + } + + pub fn iter(&self) -> ServerListIter { + ServerListIter::new(self) + } + + pub fn rotate(&self) { + self.start.fetch_add(1, Ordering::SeqCst); + } +} + +impl<'a> IntoIterator for &'a ServerList { + type Item = &'a ServerInfo; + type IntoIter = ServerListIter<'a>; + + fn into_iter(self) -> Self::IntoIter { + self.iter() + } +} + +impl ops::Deref for ServerList { + type Target = [ServerInfo]; + + fn deref(&self) -> &Self::Target { + self.servers.as_ref() + } +} + + +//------------ ServerListCounter --------------------------------------------- + +#[derive(Clone, Debug)] +pub struct ServerListCounter { + cur: usize, + end: usize, +} + +impl ServerListCounter { + fn new(list: &ServerList) -> Self { + // We modulo the start value here to prevent hick-ups towards the + // end of usize’s range. + let start = list.start.load(Ordering::Relaxed) % list.servers.len(); + ServerListCounter { + cur: start, + end: start + list.servers.len(), + } + } + + pub fn next(&mut self) { + if self.cur < self.end { + self.cur += 1 + } + } + + pub fn info<'a>(&self, list: &'a ServerList) -> Option<&'a ServerInfo> { + if self.cur == self.end { + None + } + else { + Some(&list[self.cur % list.servers.len()]) + } + } +} + + + +//------------ ServerListIter ------------------------------------------------ + +#[derive(Clone, Debug)] +pub struct ServerListIter<'a> { + servers: &'a ServerList, + counter: ServerListCounter, +} + +impl<'a> ServerListIter<'a> { + fn new(list: &'a ServerList) -> Self { + ServerListIter { + servers: list, + counter: ServerListCounter::new(list) + } + } +} + +impl<'a> Iterator for ServerListIter<'a> { + type Item = &'a ServerInfo; + + fn next(&mut self) -> Option { + self.counter.next(); + self.counter.info(self.servers) + } +} + diff --git a/domain-resolv/src/net/tcp.rs b/domain-resolv/src/net/tcp.rs new file mode 100644 index 00000000..1b6571ae --- /dev/null +++ b/domain-resolv/src/net/tcp.rs @@ -0,0 +1,101 @@ + +use std::io; +use std::net::SocketAddr; +use domain_core::bits::message::Message; +use domain_core::bits::query::{QueryMessage, StreamQueryMessage}; +use tokio::io::{read_exact, write_all, ReadExact, WriteAll}; +use tokio::net::tcp::{ConnectFuture, TcpStream}; +use tokio::prelude::{Async, Future}; +use ::resolver::Answer; +use super::util::DecoratedFuture; + +#[derive(Debug)] +pub enum TcpQuery { + Connect(DecoratedFuture), + Send(WriteAll), + RecvPrelude(DecoratedFuture, QueryMessage>), + RecvMessage(DecoratedFuture>, QueryMessage>), + Error(Option), + Done +} + +impl TcpQuery { + pub fn new( + query: QueryMessage, + addr: SocketAddr, + ) -> Self { + TcpQuery::Connect(DecoratedFuture::new( + TcpStream::connect(&addr), query + )) + } +} + +impl Future for TcpQuery { + type Item = Answer; + type Error = io::Error; + + fn poll(&mut self) -> Result, Self::Error> { + let (next, res) = match *self { + TcpQuery::Connect(ref mut fut) => { + let (sock, query) = try_ready!(fut.poll()); + ( + TcpQuery::Send(write_all(sock, query.into())), + Ok(Async::NotReady) + ) + } + TcpQuery::Send(ref mut write) => { + let (sock, query) = try_ready!(write.poll()); + ( + TcpQuery::RecvPrelude(DecoratedFuture::new( + read_exact(sock, [0u8; 2]), query.unwrap() + )), + Ok(Async::NotReady) + ) + } + TcpQuery::RecvPrelude(ref mut fut) => { + let ((sock, buf), query) = try_ready!(fut.poll()); + let len = (buf[0] as usize) << 8 | buf[1] as usize; + ( + TcpQuery::RecvMessage(DecoratedFuture::new( + read_exact(sock, vec![0; len]), query + )), + Ok(Async::NotReady) + ) + } + TcpQuery::RecvMessage(ref mut fut) => { + let ((sock, buf), query) = try_ready!(fut.poll()); + if let Ok(answer) = Message::from_bytes(buf.into()) { + if answer.is_answer(&query) { + (TcpQuery::Done, Ok(Async::Ready(answer.into()))) + } + else { + ( + TcpQuery::RecvPrelude(DecoratedFuture::new( + read_exact(sock, [0; 2]), query + )), + Ok(Async::NotReady) + ) + } + } + else { + ( + TcpQuery::Done, + Err(io::Error::new(io::ErrorKind::Other, "short buf")) + ) + } + } + TcpQuery::Error(ref mut err) => { + if let Some(err) = err.take() { + (TcpQuery::Done, Err(err)) + } + else { + panic!("polled resolved future") + } + } + TcpQuery::Done => panic!("polled resolved future"), + }; + *self = next; + res + } +} + diff --git a/domain-resolv/src/net/udp.rs b/domain-resolv/src/net/udp.rs new file mode 100644 index 00000000..ae351d3f --- /dev/null +++ b/domain-resolv/src/net/udp.rs @@ -0,0 +1,144 @@ + +use std::io; +use std::net::SocketAddr; +use domain_core::bits::query::{DgramQueryMessage, QueryMessage}; +use domain_core::bits::message::Message; +use tokio::net::udp::{RecvDgram, SendDgram, UdpSocket}; +use tokio::prelude::{Async, Future}; +use ::resolver::Answer; +use super::util::DecoratedFuture; + + +//------------ Module Configuration ------------------------------------------ + +/// How many times do we try a new random port if we get ‘address in use.’ +const RETRY_RANDOM_PORT: usize = 10; + + +//------------ UdpQuery ------------------------------------------------------ + +#[derive(Debug)] +pub enum UdpQuery { + Send { + send: SendDgram, + addr: SocketAddr, + recv_size: usize, + }, + Recv { + recv: DecoratedFuture>, QueryMessage>, + addr: SocketAddr, + recv_size: usize, + }, + Error(Option), + Done +} + +impl UdpQuery { + pub fn new( + query: QueryMessage, + addr: SocketAddr, + recv_size: usize, + ) -> Self { + let sock = match Self::bind(addr.is_ipv4()) { + Ok(sock) => sock, + Err(err) => { + return UdpQuery::Error(Some(err)) + } + }; + if let Err(err) = sock.connect(&addr) { + return UdpQuery::Error(Some(err)) + } + UdpQuery::Send { + send: sock.send_dgram(DgramQueryMessage::from(query), &addr), + addr, + recv_size + } + } + + /// Creates a bound UDP socket. + /// + /// We are supposed to pick a random local port for socket for extra + /// protection. So we try just that here. + fn bind(v4: bool) -> Result { + let mut i = 0; + loop { + let local = if v4 { ([0u8; 4], 0).into() } + else { ([0u16; 8], 0).into() }; + match UdpSocket::bind(&local) { + Ok(sock) => return Ok(sock), + Err(err) => { + if i == RETRY_RANDOM_PORT { + return Err(err); + } + else { + i += 1 + } + } + } + } + } +} + +impl Future for UdpQuery { + type Item = Answer; + type Error = io::Error; + + fn poll(&mut self) -> Result, Self::Error> { + let (next, res) = match *self { + UdpQuery::Send { ref mut send, addr, recv_size } => { + let (sock, query) = try_ready!(send.poll()); + ( + UdpQuery::Recv { + recv: DecoratedFuture::new( + sock.recv_dgram(vec![0; recv_size]), + query.unwrap() + ), + addr, recv_size + }, + Ok(Async::NotReady) + ) + } + UdpQuery::Recv { ref mut recv, addr, recv_size } => { + let ((sock, mut buf, len, recv_addr), query) + = try_ready!(recv.poll()); + buf.truncate(len); + if let Ok(answer) = Message::from_bytes(buf.into()) { + if addr == recv_addr && answer.is_answer(&query) { + (UdpQuery::Done, Ok(Async::Ready(answer.into()))) + } + else { + ( + UdpQuery::Recv { + // XXX We should reuse the buffer. + recv: DecoratedFuture::new( + sock.recv_dgram(vec![0; recv_size]), + query + ), + addr, recv_size + }, + Ok(Async::NotReady) + ) + } + } + else { + ( + UdpQuery::Done, + Err(io::Error::new(io::ErrorKind::Other, "short buf")) + ) + } + } + UdpQuery::Error(ref mut err) => { + if let Some(err) = err.take() { + (UdpQuery::Done, Err(err)) + } + else { + panic!("polled resolved future") + } + } + UdpQuery::Done => panic!("polling a resolved future"), + }; + *self = next; + res + } +} + diff --git a/domain-resolv/src/net/util.rs b/domain-resolv/src/net/util.rs new file mode 100644 index 00000000..db2bf6f2 --- /dev/null +++ b/domain-resolv/src/net/util.rs @@ -0,0 +1,33 @@ +//! Utility types for networking. + +use tokio::prelude::{Async, Future}; + + +//------------ DecoratedFuture ----------------------------------------------- + +/// A future that stores data until resolved. +/// +/// This future takes a future and some additional data. If the inner future +/// resolves successfully, the decorated future will return both the inner +/// future’s result as well as the data. +#[derive(Debug)] +pub struct DecoratedFuture(F, Option); + +impl DecoratedFuture { + pub fn new(fut: F, data: T) -> Self { + DecoratedFuture(fut, Some(data)) + } +} + +impl Future for DecoratedFuture { + type Item = (F::Item, T); + type Error = F::Error; + + fn poll(&mut self) -> Result, Self::Error> { + Ok(Async::Ready(( + try_ready!(self.0.poll()), + self.1.take().expect("polling a resolved future") + ))) + } +} + diff --git a/domain-resolv/src/resolver.rs b/domain-resolv/src/resolver.rs index 2b8493df..d6799f5f 100644 --- a/domain-resolv/src/resolver.rs +++ b/domain-resolv/src/resolver.rs @@ -8,15 +8,15 @@ use std::{io, ops}; use std::sync::Arc; -use domain_core::bits::{Message, MessageBuilder, Question}; -use domain_core::bits::message_builder::OptBuilder; +use domain_core::bits::{Message, Question}; +use domain_core::bits::query::{QueryBuilder, QueryMessage}; use domain_core::bits::name::ToDname; use domain_core::iana::Rcode; -use futures::{Future, FutureExt, TryFutureExt}; -use futures_util::compat::TokioDefaultSpawn; -use tokio::prelude::Future as TokioFuture; +use tokio::prelude::{Async, Future}; +use tokio::prelude::future::lazy; +use tokio::runtime::Runtime; use super::conf::{ResolvConf, ResolvOptions}; -use super::net::{query_server, ServerList}; +use super::net::{ServerInfo, ServerList, ServerListCounter, ServerQuery}; //------------ Resolver ------------------------------------------------------ @@ -43,7 +43,7 @@ use super::net::{query_server, ServerList}; #[derive(Clone, Debug)] pub struct Resolver(Arc); -/// The actual resolver. + #[derive(Debug)] struct ResolverInner { /// Preferred servers. @@ -54,7 +54,6 @@ struct ResolverInner { /// Resolver options. options: ResolvOptions, - } @@ -66,48 +65,35 @@ impl Resolver { /// Creates a new resolver using the given configuraiton. pub fn from_conf(conf: ResolvConf) -> Self { - Resolver(Arc::new( - ResolverInner { - preferred: ServerList::from_conf(&conf, |s| { - s.transport.is_preferred() - }), - stream: ServerList::from_conf(&conf, |s| { - s.transport.is_stream() - }), - options: conf.options - } - )) + Resolver(Arc::new(ResolverInner::from_conf(conf))) } - /// Queries the resolver for an answer to a question. - pub fn query( - &self, - question: Q - ) -> impl Future> - where N: ToDname, Q: Into> { - // 512 bytes should be enough for a domain name that is at most 255 - // bytes plus whatever EDNS there’ll be. - let mut msg = MessageBuilder::new_tcp(512); - msg.push(question).unwrap(); - - if self.0.options.recurse { - msg.header_mut().set_rd(true); - } - - let mut msg = msg.opt().unwrap(); - // Message size won’t change anymore, so we can update the prelude. - let len = msg.preview().len() - 2; - assert!(len <= usize::from(::std::u16::MAX)); - msg.prelude_mut()[0] = (len >> 8) as u8; - msg.prelude_mut()[1] = len as u8; - - async_query(self.0.clone(), msg) + pub fn options(&self) -> &ResolvOptions { + &self.0.options + } +} + +impl ResolverInner { + fn from_conf(conf: ResolvConf) -> Self { + ResolverInner { + preferred: ServerList::from_conf(&conf, |s| { + s.transport.is_preferred() + }), + stream: ServerList::from_conf(&conf, |s| { + s.transport.is_stream() + }), + options: conf.options + } } } -/// # Shortcuts -/// impl Resolver { + /// Queries the resolver for an answer to a question. + pub fn query(&self, question: Q) -> Query + where N: ToDname, Q: Into> { + Query::new(self.clone(), question) + } + /// Synchronously perform a DNS operation atop a standard resolver. /// /// This associated functions removes almost all boiler plate for the @@ -118,99 +104,217 @@ impl Resolver { /// The only argument is a closure taking a reference to a `Resolver` /// and returning a future. Whatever that future resolves to will be /// returned. - pub fn run(op: F) -> R::Output - where R: Future, F: FnOnce(&Resolver) -> R { + pub fn run(op: F) -> Result + where + R: Future + Send + 'static, + R::Item: Send + 'static, + R::Error: Send + 'static, + F: FnOnce(Resolver) -> R + Send + 'static, + { Self::run_with_conf(ResolvConf::default(), op) } - /// Synchronously perform a DNS operation atop a configuredresolver. + /// Synchronously perform a DNS operation atop a configured resolver. /// /// This is like [`run()`] but also takes a resolver configuration for /// tailor-making your own resolver. /// /// [`run()`]: #method.run - pub fn run_with_conf(conf: ResolvConf, op: F) -> R::Output - where R: Future, F: FnOnce(&Resolver) -> R { + pub fn run_with_conf( + conf: ResolvConf, + op: F + ) -> Result + where + R: Future + Send + 'static, + R::Item: Send + 'static, + R::Error: Send + 'static, + F: FnOnce(Resolver) -> R + Send + 'static, + { let resolver = Self::from_conf(conf); - op(&resolver).boxed().unit_error().compat(TokioDefaultSpawn) - .wait().unwrap() + let mut runtime = Runtime::new().unwrap(); // XXX unwrap + let res = runtime.block_on(lazy(|| op(resolver))); + runtime.shutdown_on_idle().wait().unwrap(); + res } } +//------------ Query --------------------------------------------------------- -async fn async_query( - resolver: Arc, - mut message: OptBuilder, -) -> Result { - let mut stream = false; - for _ in 0..resolver.options.attempts { - let preferred = resolver.preferred.iter(); - if resolver.options.rotate { - resolver.preferred.rotate(); +#[derive(Debug)] +pub struct Query { + /// The resolver whose configuration we are using. + resolver: Resolver, + + /// Are we still in the preferred server list or have gone streaming? + preferred: bool, + + /// The number of attempts, starting with zero. + attempt: usize, + + /// The index in the server list we currently trying. + counter: ServerListCounter, + + /// The server query we are currently performing. + /// + /// If this is an error, we had to bail out before ever starting a query. + query: Result>, + + /// The query message we currently work on. + /// + /// This is an option so we can take it out temporarily to manipulate it. + message: Option, +} + +impl Query { + fn new(resolver: Resolver, question: Q) -> Self + where N: ToDname, Q: Into> { + let message = QueryBuilder::new(question).freeze(); + let (preferred, counter) = if resolver.options().use_vc { + (false, resolver.0.stream.counter(resolver.options().rotate)) } - for server in preferred { - println!("trying {:?}", server.addr); - match await!(query_server(server, message)) { - (ret_message, Ok(answer)) => { - println!("got answer"); - if answer.is_final() { - return Ok(answer) + else { + (true, resolver.0.preferred.counter(resolver.options().rotate)) + }; + let mut res = Query { + resolver, + preferred, + attempt: 0, + counter, + query: Err(None), + message: Some(message) + }; + res.query = match res.start_query() { + Some(query) => Ok(query), + None => Err(Some(QueryError::NoServers)) + }; + res + } + + /// Starts a new query for the current server. + /// + /// Prepares the query message and then starts the server query. Returns + /// `None` if a query cannot be started because + fn start_query(&mut self) -> Option { + let mut message = self.message.take().unwrap().unfreeze(); + let (message, res) = { + let info = match self.current_server() { + Some(info) => info, + None => return None + }; + info.prepare_message(&mut message); + let message = message.freeze(); + let res = ServerQuery::new(message.clone(), info); + (message, res) + }; + self.message = Some(message); + Some(res) + } + + /// Returns the info for the current server. + fn current_server(&self) -> Option<&ServerInfo> { + let list = if self.preferred { &self.resolver.0.preferred } + else { &self.resolver.0.stream }; + self.counter.info(list) + } + + + fn switch_to_stream(&mut self) -> bool { + self.preferred = false; + self.attempt = 0; + self.counter = self.resolver.0.stream.counter( + self.resolver.options().rotate + ); + match self.start_query() { + Some(query) => { + self.query = Ok(query); + true + } + None => { + self.query = Err(None); + false + } + } + } + + fn next_server(&mut self) { + self.counter.next(); + if let Some(query) = self.start_query() { + self.query = Ok(query); + return; + } + self.attempt += 1; + if self.attempt >= self.resolver.options().attempts { + self.query = Err(Some(QueryError::GivingUp)); + return; + } + self.counter = if self.preferred { + self.resolver.0.preferred.counter(self.resolver.options().rotate) + } + else { + self.resolver.0.stream.counter(self.resolver.options().rotate) + }; + self.query = match self.start_query() { + Some(query) => Ok(query), + None => Err(Some(QueryError::GivingUp)) + } + } +} + +impl Future for Query { + type Item = Answer; + type Error = QueryError; + + fn poll(&mut self) -> Result, Self::Error> { + let answer = { + let query = match self.query { + Ok(ref mut query) => query, + Err(ref mut err) => { + let err = err.take(); + match err { + Some(err) => return Err(err), + None => panic!("polled a resolved future") } - else if answer.is_truncated() { - message = ret_message; - stream = true; - break; - } - message = ret_message; - } - (ret_message, Err(err)) => { - println!("got error {}", err); - message = ret_message; } }; - } - } - if stream { - await!(async_stream_query(resolver, message)) - } - else { - Err(QueryError::GivingUp) - } -} - -async fn async_stream_query( - resolver: Arc, - mut message: OptBuilder, -) -> Result { - for _ in 0..resolver.options.attempts { - let streams = resolver.stream.iter(); - if resolver.options.rotate { - resolver.stream.rotate(); - } - for server in streams { - match await!(query_server(server, message)) { - (ret_message, Ok(answer)) => { - if answer.is_final() { - return Ok(answer) + match query.poll() { + Ok(Async::NotReady) => return Ok(Async::NotReady), + Ok(Async::Ready(answer)) => Some(answer), + Err(_) => None, + } + }; + match answer { + Some(answer) => { + if answer.header().rcode() == Rcode::FormErr + && self.current_server().unwrap().does_edns() + { + // FORMERR with EDNS: turn off EDNS and try again. + self.current_server().unwrap().disable_edns(); + self.query = Ok(self.start_query().unwrap()); + } + else if answer.header().rcode() == Rcode::ServFail { + // SERVFAIL: go to next server. + self.next_server(); + } + else if answer.header().tc() && self.preferred + && !self.resolver.options().ign_tc + { + // Truncated. If we can, switch to stream transports. + if !self.switch_to_stream() { + return Ok(Async::Ready(answer)) } - message = ret_message; } - (ret_message, Err(_)) => { - message = ret_message; + else { + // I guess we have an answer ... + self.query = Err(None); // Make it panic if polled again. + return Ok(Async::Ready(answer)); } - }; + } + None => { + self.next_server(); + } } - } - Err(QueryError::GivingUp) -} - - -//--- Default - -impl Default for Resolver { - fn default() -> Self { - Self::new() + self.poll() } } @@ -269,6 +373,19 @@ impl AsRef for Answer { #[derive(Debug)] pub enum QueryError { + NoServers, GivingUp, + MalformedAnswer, // XXX Return this for when parsing fails. Io(io::Error) } + +impl QueryError { + pub fn merge(self, other: Self) -> Self { + if let QueryError::GivingUp = self { + other + } + else { + self + } + } +} diff --git a/domain-resolv/src/search.rs b/domain-resolv/src/search.rs new file mode 100644 index 00000000..30ac03ba --- /dev/null +++ b/domain-resolv/src/search.rs @@ -0,0 +1,195 @@ + +use std::ops; +use domain_core::bits::name::{Chain, Dname, ToRelativeDname}; +use tokio::prelude::{Async, Future, Poll}; +use ::resolver::Resolver; + +//------------ search -------------------------------------------------------- + +pub fn search( + resolver: &Resolver, + name: N, + op: F +) -> Search +where + N: ToRelativeDname + Clone, + F: Fn(&Resolver, Chain) -> R, + R: Future +{ + Search::new(resolver, name, op) +} + + +//------------ SearchList ---------------------------------------------------- + +#[derive(Clone, Debug, Default)] +pub struct SearchList { + search: Vec, +} + +impl SearchList { + pub fn new() -> Self { + Self::default() + } + + pub fn push(&mut self, name: Dname) { + if !name.is_root() && !self.search.contains(&name) { + self.search.push(name) + } + } + + pub fn as_slice(&self) -> &[Dname] { + self.as_ref() + } +} + +impl From for SearchList { + fn from(name: Dname) -> Self { + let mut res = Self::new(); + res.push(name); + res + } +} + + +//--- AsRef and Deref + +impl AsRef<[Dname]> for SearchList { + fn as_ref(&self) -> &[Dname] { + self.search.as_ref() + } +} + +impl ops::Deref for SearchList { + type Target = [Dname]; + + fn deref(&self) -> &Self::Target { + self.as_ref() + } +} + + +//------------ SearchIter ---------------------------------------------------- + +/// An iterator gained from applying a search list to a domain name. +/// +/// The iterator represents how a resolver attempts to derive an absolute +/// domain name from a relative name. +/// +/// For this purpose, the resolver’s configuration contains a search list, +/// a list of absolute domain names that are appened in turn to the domain +/// name. In addition, if the name contains enough dots (specifically, +/// `ResolvConf::ndots` which defaults to just one) it is first tried as if +/// it were an absolute by appending the root labels. +#[derive(Clone, Debug)] +pub struct SearchIter { + name: N, + resolver: Resolver, + pos: Option, +} + +impl SearchIter { + pub fn new(resolver: &Resolver, name: N) -> Self { + SearchIter { + name, + resolver: resolver.clone(), + pos: Some(0), + } + } +} + +impl Iterator for SearchIter { + type Item = Chain; + + fn next(&mut self) -> Option { + while let Some(pos) = self.pos { + if pos >= self.resolver.options().search.len() { + self.pos = None; + return Some(self.name.clone().chain_root()) + } + else { + self.pos = Some(pos + 1); + let name = self.name.clone() + .chain(self.resolver.options().search[pos].clone()); + if let Ok(name) = name { + return Some(name) + } + } + } + None + } +} + + +//------------ SearchFuture -------------------------------------------------- + +#[derive(Debug)] +pub struct Search +where + N: ToRelativeDname + Clone, + F: Fn(&Resolver, Chain) -> R, + R: Future +{ + iter: SearchIter, + op: F, + pending: Option, +} + +impl Search +where + N: ToRelativeDname + Clone, + F: Fn(&Resolver, Chain) -> R, + R: Future +{ + fn new(resolver: &Resolver, name: N, op: F) -> Self { + let mut iter = SearchIter::new(resolver, name); + match iter.next() { + Some(name) => { + Search { + iter, + pending: Some(op(resolver, name)), + op + } + } + None => { + Search { + iter, + op, + pending: None + } + } + } + } +} + +impl Future for Search +where + N: ToRelativeDname + Clone, + F: Fn(&Resolver, Chain) -> R, + R: Future +{ + type Item = R::Item; + type Error = R::Error; + + fn poll(&mut self) -> Poll { + let err = match self.pending { + Some(ref mut pending) => match pending.poll() { + Ok(Async::NotReady) => return Ok(Async::NotReady), + Ok(Async::Ready(res)) => return Ok(Async::Ready(res)), + Err(err) => err + } + None => panic!("polled a resolved future"), + }; + + match self.iter.next() { + Some(name) => { + self.pending = Some((self.op)(&self.iter.resolver, name)); + self.poll() + } + None => { + Err(err) + } + } + } +} +