kapiti 0.0.3

The Kapiti DNS Server
Documentation
use std::convert::TryFrom;
use std::net::{SocketAddr, TcpStream};
use std::time::Duration;

use anyhow::{bail, Context, Result};
use async_native_tls_alpn::{Protocol, TlsConnector, TlsStream};
use async_trait::async_trait;
use byteorder::{BigEndian, ByteOrder};
use bytes::{BufMut, BytesMut};
use futures_lite::{AsyncReadExt, AsyncWriteExt};
use rand::Rng;
use smol::Async;
use tracing::debug;

use crate::client::DnsClient;
use crate::codec::{decoder::DNSMessageDecoder, message};
use crate::specs::message::Message;
use crate::timeout;

/// TCP size header is 16 bits, so max theoretical size is 64k
static MAX_TCP_BYTES: u16 = 65535;

pub struct Client {
    dns_server: SocketAddr,
    conn: Option<TlsStream<Async<TcpStream>>>,
    response_buffer: BytesMut,
    timeout: Duration,
}

/// DNS Client that queries a server over TLS
impl Client {
    /// Constructs a new `Client` that will query the specified `dns_server`.
    pub fn new(dns_server: SocketAddr, timeout: Duration) -> Self {
        Client {
            dns_server,
            conn: None,
            response_buffer: BytesMut::with_capacity(MAX_TCP_BYTES as usize),
            timeout,
        }
    }

    async fn connect(&mut self) -> Result<()> {
        let stream = timeout::timeout(
            Async::<TcpStream>::connect(self.dns_server.clone()),
            &self.timeout,
        )
        .await
        .context("TLS TCP connect timed out")?
        .context("TLS TCP connect failed")?;
        // Min protocol: If things don't have at least TLS1.2 by now, we should just name and shame.
        let connector = TlsConnector::new().min_protocol_version(Some(Protocol::Tlsv12));
        let stream = timeout::timeout(
            connector.connect(format!("{}", self.dns_server.ip()), stream),
            &self.timeout,
        )
        .await
        .context("TLS session connect timed out")?
        .context("TLS session connect failed")?;
        self.conn = Some(stream);
        Ok(())
    }
}

#[async_trait]
impl DnsClient for Client {
    async fn query(
        &mut self,
        request: &Message,
        query_buffer: &mut BytesMut,
    ) -> Result<Option<Message>> {
        // Reserve 2 bytes for the TLS-specific length prefix
        query_buffer.reserve(2);
        query_buffer.put_u16(0);

        // Encode the message, along with any necessary padding
        super::add_request_padding(request, MAX_TCP_BYTES, query_buffer)?;

        // Insert the resulting encoded size of the message
        // into those leading two bytes that we'd reserved.
        let message_len = u16::try_from(query_buffer.len() - 2).with_context(|| {
            format!(
                "Encoded request size {} exceeds {} limit: {}",
                query_buffer.len() - 2,
                MAX_TCP_BYTES,
                request
            )
        })?;

        query_buffer[0] = ((message_len & 0xFF00) >> 8) as u8;
        query_buffer[1] = (message_len & 0xFF) as u8;

        // Query is constructed, now let's do the request.
        if self.conn.is_none() {
            self.connect().await.with_context(|| {
                format!("Failed to connect with TLS upstream {:?}", self.dns_server)
            })?;
        }

        let request_id = rand::thread_rng().gen::<u16>();
        // For TCP, the size header means that the message actually starts at byte 2
        message::update_message_id(request_id, query_buffer, 2)?;

        debug!(
            "Raw request to {:?} ({}b): {:02X?}",
            self.dns_server,
            query_buffer.len(),
            &query_buffer[..]
        );

        // NOTE: async is useful here, since it allows us to ensure that the entire write completes within the timeout.
        // If we used sync APIs, we would risk a malicious upstream slowly allowing one byte at a time. (sync write_all loops over writes internally)
        match timeout::timeout(
            self.conn
                .as_mut()
                .expect("missing connection")
                .write_all(query_buffer.as_ref()),
            &self.timeout,
        )
        .await
        {
            Some(Ok(())) => {}
            Some(Err(e)) => {
                bail!(
                    "Failed to write to TLS upstream {:?}: {}",
                    self.dns_server,
                    e
                )
            }
            None => {
                // Mark connection as dead, reconnect again on next query
                self.conn = None;
                bail!("Timed out writing to TLS upstream {:?}", self.dns_server)
            }
        }

        // Read first two bytes to get expected response size
        let mut response_size_bytes: [u8; 2] = [0, 0];
        match timeout::timeout(
            self.conn
                .as_mut()
                .expect("missing connection")
                .read_exact(&mut response_size_bytes),
            &self.timeout,
        )
        .await
        {
            Some(Ok(())) => {}
            Some(Err(e)) => {
                bail!(
                    "Failed to read header from TLS upstream {:?}: {}",
                    self.dns_server,
                    e
                )
            }
            None => {
                // Mark connection as dead, reconnect again on next query
                self.conn = None;
                bail!(
                    "Timed out reading header from TLS upstream {:?}",
                    self.dns_server
                )
            }
        }
        // big endian
        let response_size = BigEndian::read_u16(&response_size_bytes);

        // Read remaining bytes to get response
        self.response_buffer.resize(
            usize::try_from(response_size).with_context(|| "couldn't convert u16 to usize")?,
            0,
        );
        // NOTE: async is useful here, since it allows us to ensure that the entire read completes within the timeout.
        // If we used sync APIs, we would risk a malicious upstream slowly allowing one byte at a time. (sync read_all loops over writes internally)
        match timeout::timeout(
            self.conn
                .as_mut()
                .expect("missing connection")
                .read_exact(&mut self.response_buffer),
            &self.timeout,
        )
        .await
        {
            Some(Ok(())) => {}
            Some(Err(e)) => {
                bail!(
                    "Failed to read payload from TLS upstream {:?}: {}",
                    self.dns_server,
                    e
                )
            }
            None => {
                // Mark connection as dead, reconnect again on next query
                self.conn = None;
                bail!(
                    "Timed out reading payload from TLS upstream {:?}",
                    self.dns_server
                )
            }
        }

        debug!(
            "Raw response from {:?} ({}b): {:02X?}",
            self.dns_server,
            self.response_buffer.len(),
            &self.response_buffer[..]
        );

        match DNSMessageDecoder::new().decode(&self.response_buffer[..]) {
            Ok(Some(mut response)) => {
                debug!(
                    "Untrimmed response from {:?}: {}",
                    self.dns_server, 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.
                // The PADDING is common/best practice for DoT servers.
                super::remove_response_padding(&mut response);

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