use crate::bail;
use crate::clients::mime::content_type_equal;
use crate::clients::stats::StatsBuilder;
use crate::clients::AsyncExchanger;
use crate::clients::ToUrls;
use crate::Message;
use async_trait::async_trait;
use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _};
use http::header::*;
use http::{Method, Request};
use http_body_util::{BodyExt, Full};
use hyper::body::Bytes;
use hyper_rustls::HttpsConnectorBuilder;
use hyper_util::client::legacy::connect::HttpInfo;
use hyper_util::client::legacy::Client as HyperClient;
use hyper_util::rt::TokioExecutor;
use std::net::IpAddr;
use std::net::Ipv4Addr;
use std::net::SocketAddr;
use std::time::Duration;
use url::Url;
pub const GOOGLE: &str = "https://dns.google/dns-query";
const CONTENT_TYPE_APPLICATION_DNS_MESSAGE: &str = "application/dns-message";
const DNS_QUERY_PARAM: &str = "dns";
pub struct Client {
servers: Vec<Url>,
method: Method, }
impl Default for Client {
fn default() -> Self {
Client {
servers: Vec::default(),
method: Method::GET,
}
}
}
impl Client {
pub fn new<A: ToUrls>(servers: A, method: Method) -> Result<Self, crate::Error> {
match method {
Method::GET | Method::POST => (), _ => bail!(InvalidInput, "only GET and POST allowed"),
}
let servers: Vec<_> = servers.to_urls()?.collect();
if servers.is_empty() {
return Err(crate::Error::InvalidArgument(
"at least one DoH server is required".to_string(),
));
}
if servers.iter().any(|server| server.scheme() != "https") {
return Err(crate::Error::InvalidArgument(
"DoH servers must use HTTPS".to_string(),
));
}
Ok(Self { servers, method })
}
}
#[async_trait]
impl AsyncExchanger for Client {
async fn exchange(&self, query: &Message) -> Result<Message, crate::Error> {
let mut query = query.clone();
query.id = 0;
let p = query.to_vec()?;
let https = HttpsConnectorBuilder::new()
.with_webpki_roots()
.https_only()
.enable_http1()
.enable_http2()
.build();
let client: HyperClient<_, Full<Bytes>> = HyperClient::builder(TokioExecutor::new())
.pool_idle_timeout(Duration::from_secs(30))
.http2_only(true) .build(https);
let req = Request::builder()
.method(&self.method)
.header(ACCEPT, CONTENT_TYPE_APPLICATION_DNS_MESSAGE);
let req = match self.method {
Method::GET => {
let mut buf = String::new();
URL_SAFE_NO_PAD.encode_string(p, &mut buf);
let mut url = self.servers[0].clone(); url.query_pairs_mut().append_pair(DNS_QUERY_PARAM, &buf);
let uri: http::Uri = url.as_str().parse()?;
req.uri(uri).body(Full::new(Bytes::new()))?
}
Method::POST => {
req.uri(self.servers[0].as_str()) .header(CONTENT_TYPE, CONTENT_TYPE_APPLICATION_DNS_MESSAGE)
.body(Full::new(Bytes::from(p)))? }
_ => bail!(InvalidInput, "only GET and POST allowed"),
};
let stats = StatsBuilder::start(0);
let resp = client.request(req).await?;
if let Some(content_type) = resp.headers().get(CONTENT_TYPE) {
if !content_type_equal(content_type, CONTENT_TYPE_APPLICATION_DNS_MESSAGE) {
bail!(
InvalidData,
"recevied invalid content-type: {:?} expected {}",
content_type,
CONTENT_TYPE_APPLICATION_DNS_MESSAGE,
);
}
}
if resp.status().is_success() {
let remote_addr = match resp.extensions().get::<HttpInfo>() {
Some(http_info) => http_info.remote_addr(),
None => SocketAddr::new(IpAddr::V4(Ipv4Addr::new(0, 0, 0, 0)), 0), };
let body = resp.into_body().collect().await?.to_bytes();
let mut m = Message::from_slice(&body)?;
m.stats = Some(stats.end(remote_addr, body.len()));
return Ok(m);
}
bail!(
InvalidInput,
"recevied unexpected HTTP status code: {:}",
resp.status()
);
}
}
#[cfg(test)]
mod tests {
use super::Client;
use http::Method;
#[test]
fn rejects_plaintext_http_server() {
assert!(Client::new("http://dns.example/dns-query", Method::GET).is_err());
}
}