use crate::bail;
use crate::clients::mime::content_type_equal;
use crate::clients::stats::StatsBuilder;
use crate::clients::validate_http_status;
use crate::clients::AsyncExchanger;
use crate::clients::ToUrls;
use crate::clients::{new_http_client, BoxError, HttpClient};
use crate::limits::MAX_DNS_MESSAGE_LEN;
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, Limited};
use hyper::body::Bytes;
use hyper_util::client::legacy::connect::HttpInfo;
use std::io;
use std::net::IpAddr;
use std::net::Ipv4Addr;
use std::net::SocketAddr;
use url::Url;
const MAX_DOH_BODY_SIZE: usize = MAX_DNS_MESSAGE_LEN;
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,
http_client: HttpClient,
}
impl std::panic::RefUnwindSafe for Client {}
impl std::panic::UnwindSafe for Client {}
impl Default for Client {
fn default() -> Self {
Client {
servers: Vec::default(),
method: Method::GET,
http_client: new_http_client(),
}
}
}
impl Client {
pub fn new<A: ToUrls>(servers: A, method: Method) -> Result<Self, crate::Error> {
match method {
Method::GET | Method::POST => (), _ => {
return Err(crate::Error::InvalidArgument(
"only GET and POST allowed".to_string(),
));
}
}
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,
http_client: new_http_client(),
})
}
}
#[async_trait]
impl AsyncExchanger for Client {
async fn exchange(&self, query: &Message) -> Result<Message, crate::Error> {
let server = self.servers.first().ok_or_else(|| {
crate::Error::InvalidArgument("at least one DoH server is required".to_string())
})?;
let mut query = query.clone();
query.id = 0;
let p = query.to_vec()?;
let dns_request_len = p.len();
let client = &self.http_client;
let req = Request::builder()
.method(&self.method)
.header(ACCEPT, CONTENT_TYPE_APPLICATION_DNS_MESSAGE);
let mut request_target = server.to_string();
let req = match self.method {
Method::GET => {
let mut buf = String::new();
URL_SAFE_NO_PAD.encode_string(p, &mut buf);
let mut url = server.clone(); url.query_pairs_mut().append_pair(DNS_QUERY_PARAM, &buf);
request_target = url.to_string();
let uri: http::Uri = url.as_str().parse()?;
req.uri(uri).body(
Full::new(Bytes::new())
.map_err(|error: std::convert::Infallible| -> BoxError { match error {} })
.boxed(),
)?
}
Method::POST => {
req.uri(server.as_str()) .header(CONTENT_TYPE, CONTENT_TYPE_APPLICATION_DNS_MESSAGE)
.body(
Full::new(Bytes::from(p))
.map_err(|error: std::convert::Infallible| -> BoxError {
match error {}
})
.boxed(),
)? }
_ => {
return Err(crate::Error::InvalidArgument(
"only GET and POST allowed".to_string(),
));
}
};
let stats = StatsBuilder::start(0);
log::trace!(
"DoH sending {} request to {request_target} with {dns_request_len} DNS bytes",
self.method
);
let resp = client.request(req).await?;
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), };
log::trace!("DoH remote address: {remote_addr}");
log::trace!("DoH HTTP status: {}", resp.status());
let content_type = resp.headers().get(CONTENT_TYPE).ok_or_else(|| {
io::Error::new(
io::ErrorKind::InvalidData,
"response is missing content-type",
)
})?;
log::trace!("DoH response content-type: {:?}", 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,
);
}
validate_http_status(resp.status())?;
let body = Limited::new(resp.into_body(), MAX_DOH_BODY_SIZE)
.collect()
.await
.map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?
.to_bytes();
log::trace!(
"DoH received {} DNS body bytes from {remote_addr}",
body.len()
);
let mut m = Message::from_slice(&body)?;
m.stats = Some(stats.end(remote_addr, body.len()));
return Ok(m);
}
}
#[cfg(test)]
mod tests {
use super::{Client, MAX_DOH_BODY_SIZE};
use crate::clients::validate_http_status;
use http::Method;
use http::StatusCode;
use http_body_util::{BodyExt, Full, Limited};
use hyper::body::Bytes;
#[test]
fn rejects_plaintext_http_server() {
assert!(Client::new("http://dns.example/dns-query", Method::GET).is_err());
}
#[tokio::test]
async fn rejects_oversized_response_body() {
let body = Limited::new(
Full::new(Bytes::from(vec![0; MAX_DOH_BODY_SIZE + 1])),
MAX_DOH_BODY_SIZE,
)
.collect()
.await;
assert!(body.is_err());
}
#[test]
fn validates_success_client_statuses() {
assert!(validate_http_status(StatusCode::OK).is_ok());
assert!(validate_http_status(StatusCode::BAD_REQUEST).is_err());
assert!(validate_http_status(StatusCode::INTERNAL_SERVER_ERROR).is_err());
}
}