use std::io::{self, Write};
use std::time::Duration;
use anyhow::{bail, Context, Result};
use async_trait::async_trait;
use bytes::{BufMut, BytesMut};
use hyper::header;
use hyper::{Body, Client as HttpClient, Method, Uri};
use rand::Rng;
use scopeguard;
use tracing::debug;
use crate::client::DnsClient;
use crate::codec::{decoder::DNSMessageDecoder, message};
use crate::fetcher::Fetcher;
use crate::hyper_smol;
use crate::resolver;
use crate::specs::message::Message;
use crate::timeout;
static MAX_HTTP_BYTES: u16 = 65535;
pub struct Client {
server_url: Uri,
fetcher: Fetcher,
client: HttpClient<hyper_smol::SmolConnector>,
timeout: Duration,
response_buffer: BytesMut,
}
impl Client {
pub fn new_hostname(
server_url: Uri,
resolver: resolver::Resolver,
timeout: Duration,
) -> Result<Self> {
Ok(Client {
server_url,
fetcher: Fetcher::new(
MAX_HTTP_BYTES as usize,
Some("application/dns-message".to_string()),
)
.use_http_2(),
client: hyper_smol::client_kapiti(resolver, true, false, 4096, timeout.clone()),
timeout,
response_buffer: BytesMut::with_capacity(MAX_HTTP_BYTES as usize),
})
}
pub fn new_ip(server_url: Uri, timeout: Duration) -> Result<Self> {
Ok(Client {
server_url,
fetcher: Fetcher::new(
MAX_HTTP_BYTES as usize,
Some("application/dns-message".to_string()),
)
.use_http_2(),
client: hyper_smol::client_iponly(true, timeout.clone()),
timeout,
response_buffer: BytesMut::with_capacity(MAX_HTTP_BYTES as usize),
})
}
}
#[async_trait]
impl DnsClient for Client {
async fn query(
&mut self,
request: &Message,
query_buffer: &mut BytesMut,
) -> Result<Option<Message>> {
super::add_request_padding(request, MAX_HTTP_BYTES, query_buffer)?;
let request_id = rand::thread_rng().gen::<u16>();
message::update_message_id(request_id, query_buffer, 0)?;
let mut response_buffer = scopeguard::guard(&mut self.response_buffer, |buf| {
buf.clear();
});
debug!(
"Raw request to {} ({}b): {:02X?}",
self.server_url,
query_buffer.len(),
&query_buffer[..]
);
let body = hyper::body::Bytes::from(query_buffer.clone());
let request = self
.fetcher
.request_builder(&Method::POST, &self.server_url)
.header(header::CONTENT_TYPE, "application/dns-message")
.header(header::CONTENT_LENGTH, body.len())
.body(Body::from(body))
.context("Failed to build DoH request")?;
let mut response = timeout::timeout(self.client.request(request), &self.timeout)
.await
.with_context(|| format!("Timed out querying DoH upstream {:?}", self.server_url))?
.with_context(|| format!("Failed to query DoH upstream {:?}", self.server_url))?;
if !response.status().is_success() {
bail!(
"HTTP POST to {} returned status: {}",
self.server_url,
response.status()
);
}
{
let mut writer = BytesWriter::new(&mut response_buffer);
self.fetcher
.write_response(&self.server_url.to_string(), &mut writer, &mut response)
.await?;
}
debug!(
"Raw response from {} ({}b): {:02X?}",
self.server_url,
response_buffer.len(),
&response_buffer[..]
);
match DNSMessageDecoder::new().decode(&response_buffer[..]) {
Ok(Some(mut response)) => {
debug!("Response from {}: {}", self.server_url, response);
if response.header.truncated {
return Ok(None);
}
if response.header.id != request_id {
bail!(
"Returned transaction id {:?} doesn't match sent {:?}",
response.header.id,
request_id
);
}
super::remove_response_padding(&mut response);
Ok(Some(response))
}
Ok(None) => {
debug!(
"Unable to parse response from server={} to request={:02X?}: {:02X?}",
self.server_url,
&query_buffer[..],
&response_buffer[..],
);
Ok(None)
}
Err(e) => {
Err(e).context(format!(
"Failed to parse response from server={} to request={:02X?}: {:02X?}",
self.server_url,
&query_buffer[..],
&response_buffer[..],
))
}
}
}
}
struct BytesWriter<'a> {
inner: &'a mut BytesMut,
}
impl<'a> BytesWriter<'a> {
fn new(inner: &'a mut BytesMut) -> BytesWriter<'a> {
BytesWriter { inner }
}
}
impl<'a> Write for BytesWriter<'a> {
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
if self.inner.remaining_mut() >= buf.len() {
self.inner.put_slice(buf);
Ok(buf.len())
} else {
Err(io::Error::new(
io::ErrorKind::InvalidInput,
format!(
"Unable to write {} bytes into buffer: {}/{} remaining",
buf.len(),
self.inner.remaining_mut(),
self.inner.len()
),
))
}
}
fn flush(&mut self) -> io::Result<()> {
Ok(())
}
}