use std::net::SocketAddr;
use std::sync::Arc;
use std::time::Duration;
use async_trait::async_trait;
use dns_lattice_core::{Error, Result};
use dns_lattice_model::Message;
use rustls_pki_types::ServerName;
use tokio::net::TcpStream;
use tokio::time::timeout;
use tokio_rustls::TlsConnector;
use tokio_rustls::rustls::ClientConfig;
use super::{UpstreamBackend, framed_query};
#[derive(Clone)]
pub struct DotBackendConfig {
pub server: SocketAddr,
pub server_name: ServerName<'static>,
pub tls_config: Arc<ClientConfig>,
pub connect_timeout: Duration,
pub read_timeout: Duration,
}
impl DotBackendConfig {
pub fn with_webpki_roots(
server: SocketAddr,
server_name: ServerName<'static>,
connect_timeout: Duration,
read_timeout: Duration,
) -> Self {
let mut root_store = tokio_rustls::rustls::RootCertStore::empty();
root_store.extend(webpki_roots::TLS_SERVER_ROOTS.iter().cloned());
let tls_config = ClientConfig::builder()
.with_root_certificates(root_store)
.with_no_client_auth();
Self {
server,
server_name,
tls_config: Arc::new(tls_config),
connect_timeout,
read_timeout,
}
}
}
pub struct DotBackend {
config: DotBackendConfig,
}
impl DotBackend {
pub fn new(config: DotBackendConfig) -> Self {
Self { config }
}
}
#[async_trait]
impl UpstreamBackend for DotBackend {
async fn resolve(&self, query: &Message) -> Result<Message> {
let stream = timeout(
self.config.connect_timeout,
TcpStream::connect(self.config.server),
)
.await
.map_err(|_| Error::Timeout)?
.map_err(|err| Error::Transport(err.to_string()))?;
let connector = TlsConnector::from(self.config.tls_config.clone());
let mut tls_stream = timeout(
self.config.read_timeout,
connector.connect(self.config.server_name.clone(), stream),
)
.await
.map_err(|_| Error::Timeout)?
.map_err(map_tls_connect_error)?;
framed_query(&mut tls_stream, self.config.read_timeout, query).await
}
}
fn map_tls_connect_error(err: std::io::Error) -> Error {
if error_chain_is_tls(&err) {
Error::Tls(err.to_string())
} else {
Error::Transport(err.to_string())
}
}
fn error_chain_is_tls(err: &(dyn std::error::Error + 'static)) -> bool {
let mut current: Option<&(dyn std::error::Error + 'static)> = Some(err);
while let Some(node) = current {
if node.downcast_ref::<tokio_rustls::rustls::Error>().is_some() {
return true;
}
if let Some(io_err) = node.downcast_ref::<std::io::Error>()
&& let Some(inner) = io_err.get_ref()
&& error_chain_is_tls(inner)
{
return true;
}
current = node.source();
}
false
}
#[cfg(test)]
mod tests {
use super::*;
use dns_lattice_model::{Class, Header, Name, Opcode, Question, Rcode, RecordType};
use rcgen::{CertifiedKey, generate_simple_self_signed};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpListener;
use tokio_rustls::TlsAcceptor;
use tokio_rustls::rustls::pki_types::{CertificateDer, PrivateKeyDer};
use tokio_rustls::rustls::{RootCertStore, ServerConfig};
fn query_for(name: &str) -> Message {
Message {
header: Header {
id: 11,
qr: false,
opcode: Opcode::Query,
authoritative: false,
truncated: false,
recursion_desired: true,
recursion_available: false,
rcode: Rcode::NoError,
},
questions: vec![Question {
name: Name::from_ascii(name).unwrap(),
qtype: RecordType::A,
qclass: Class::In,
}],
answers: vec![],
authorities: vec![],
additionals: vec![],
}
}
fn answer_for(name: &str, id: u16) -> Message {
let mut msg = query_for(name);
msg.header.id = id;
msg.header.qr = true;
msg
}
fn self_signed_fixture() -> (ServerConfig, ClientConfig, ServerName<'static>) {
let CertifiedKey { cert, signing_key } =
generate_simple_self_signed(vec!["localhost".to_string()]).unwrap();
let cert_der: CertificateDer<'static> = cert.der().clone();
let key_der: PrivateKeyDer<'static> =
PrivateKeyDer::try_from(signing_key.serialize_der()).unwrap();
let server_config = ServerConfig::builder()
.with_no_client_auth()
.with_single_cert(vec![cert_der.clone()], key_der)
.unwrap();
let mut roots = RootCertStore::empty();
roots.add(cert_der).unwrap();
let client_config = ClientConfig::builder()
.with_root_certificates(roots)
.with_no_client_auth();
let server_name = ServerName::try_from("localhost").unwrap();
(server_config, client_config, server_name)
}
#[tokio::test]
async fn dot_backend_resolves_against_a_loopback_tls_server() {
let (server_config, client_config, server_name) = self_signed_fixture();
let acceptor = TlsAcceptor::from(Arc::new(server_config));
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let responder = tokio::spawn(async move {
let (tcp_stream, _) = listener.accept().await.unwrap();
let mut tls_stream = acceptor.accept(tcp_stream).await.unwrap();
let mut len_buf = [0u8; 2];
tls_stream.read_exact(&mut len_buf).await.unwrap();
let len = u16::from_be_bytes(len_buf) as usize;
let mut payload = vec![0u8; len];
tls_stream.read_exact(&mut payload).await.unwrap();
let query = Message::decode(&payload).unwrap();
let response = answer_for("example.com", query.header.id);
let bytes = response.encode().unwrap();
let framed_len: u16 = bytes.len().try_into().unwrap();
let mut framed = Vec::new();
framed.extend_from_slice(&framed_len.to_be_bytes());
framed.extend_from_slice(&bytes);
tls_stream.write_all(&framed).await.unwrap();
});
let backend = DotBackend::new(DotBackendConfig {
server: addr,
server_name,
tls_config: Arc::new(client_config),
connect_timeout: Duration::from_secs(2),
read_timeout: Duration::from_secs(2),
});
let answer = backend
.resolve(&query_for("example.com"))
.await
.expect("dot backend resolves");
assert!(answer.header.qr);
responder.await.unwrap();
}
#[tokio::test]
async fn dot_backend_returns_tls_error_on_untrusted_certificate() {
let (server_config, _matching_client_config, _server_name) = self_signed_fixture();
let (_other_server_config, untrusting_client_config, server_name) = self_signed_fixture();
let acceptor = TlsAcceptor::from(Arc::new(server_config));
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let responder = tokio::spawn(async move {
let (tcp_stream, _) = listener.accept().await.unwrap();
let _ = acceptor.accept(tcp_stream).await;
});
let backend = DotBackend::new(DotBackendConfig {
server: addr,
server_name,
tls_config: Arc::new(untrusting_client_config),
connect_timeout: Duration::from_secs(2),
read_timeout: Duration::from_secs(2),
});
let err = backend
.resolve(&query_for("example.com"))
.await
.expect_err("untrusted certificate fails the tls handshake");
assert!(matches!(err, Error::Tls(_)));
let _ = responder.await;
}
#[tokio::test]
async fn dot_backend_returns_transport_when_peer_closes_before_tls() {
let (_server_config, client_config, server_name) = self_signed_fixture();
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let responder = tokio::spawn(async move {
let (stream, _) = listener.accept().await.unwrap();
drop(stream);
});
let backend = DotBackend::new(DotBackendConfig {
server: addr,
server_name,
tls_config: Arc::new(client_config),
connect_timeout: Duration::from_secs(2),
read_timeout: Duration::from_secs(2),
});
let err = backend
.resolve(&query_for("example.com"))
.await
.expect_err("a peer that closes before TLS is a transport failure");
assert!(matches!(err, Error::Transport(_)));
responder.await.unwrap();
}
#[tokio::test]
async fn dot_backend_times_out_when_server_never_completes_handshake() {
let (server_config, client_config, server_name) = self_signed_fixture();
let _acceptor = TlsAcceptor::from(Arc::new(server_config));
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let (_never_sent, wait_forever) = tokio::sync::oneshot::channel::<()>();
let responder = tokio::spawn(async move {
let (_tcp_stream, _) = listener.accept().await.unwrap();
let _ = wait_forever.await;
});
let backend = DotBackend::new(DotBackendConfig {
server: addr,
server_name,
tls_config: Arc::new(client_config),
connect_timeout: Duration::from_secs(2),
read_timeout: Duration::from_millis(50),
});
let err = backend
.resolve(&query_for("example.com"))
.await
.expect_err("handshake does not complete within the timeout budget");
assert_eq!(err, Error::Timeout);
responder.abort();
}
}