kapiti 0.0.3

The Kapiti DNS Server
Documentation
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,
}

/// DNS Client that queries a server over HTTPS (DoH/RFC8484)
impl Client {
    /// Constructs a new `Client` that will query the specified `server_url`.
    /// Uses the provided bootstrap `resolver` to resolve `server_url` itself.
    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()),
            )
            // Note that hyper will reject requests with "request has unsupported HTTP version",
            // unless we ALSO set "http2_only(true)" in the Client builder.
            .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),
        })
    }

    /// Constructs a new `Client` that will query the specified `server_url`, which must be for an IP.
    /// The full URL is provided instead of a SocketAddr so that we retain the URL path.
    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()),
            )
            // Note that hyper will reject requests with "request has unsupported HTTP version",
            // unless we ALSO set "http2_only(true)" in the Client builder.
            .use_http_2(),
            // Use a client that will refuse to resolve hostnames - expecting only IPs
            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>> {
        // Encode the message, along with any necessary padding
        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)?;

        // Ensure that response_buffer size is reset when we're done with it,
        // regardless of success or error. The socket read is based on len, not capacity.
        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[..]
        );

        // Hyper apparently requires a static copy of the body content?
        // No idea why, and don't care to find out...
        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()
            );
        }

        {
            // Write response payload into response_buffer
            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 {
                    // Message claims to be truncated, shouldn't happen but let's bail anyway
                    return Ok(None);
                }
                if response.header.id != request_id {
                    bail!(
                        "Returned transaction id {:?} doesn't match sent {:?}",
                        response.header.id,
                        request_id
                    );
                }

                // Filter out any EDNS PADDING from the response before returning it.
                // This PADDING is less common for DoH than it is for DoT, but doesn't hurt to check.
                super::remove_response_padding(&mut response);

                Ok(Some(response))
            }
            Ok(None) => {
                // Message was likely corrupted somehow, despite us receiving all the data in the payload
                debug!(
                    "Unable to parse response from server={} to request={:02X?}: {:02X?}",
                    self.server_url,
                    &query_buffer[..],
                    &response_buffer[..],
                );
                Ok(None)
            }
            Err(e) => {
                // Other parse error
                Err(e).context(format!(
                    "Failed to parse response from server={} to request={:02X?}: {:02X?}",
                    self.server_url,
                    &query_buffer[..],
                    &response_buffer[..],
                ))
            }
        }
    }
}

/// Pass-through writer that counts the number of bytes that have been written.
/// Used to consistently measure the decompressed size of a download.
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(())
    }
}