use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr};
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use std::time::Duration;
use async_trait::async_trait;
use dns_lattice_core::{Error, Result};
use dns_lattice_model::Message;
use quinn::crypto::rustls::QuicClientConfig;
use quinn::{ClientConfig, ConnectionError, Endpoint, RecvStream, SendStream};
use rustls::ClientConfig as RustlsClientConfig;
use rustls_pki_types::ServerName;
use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
use tokio::time::timeout;
use super::{UpstreamBackend, framed_query};
const DOQ_ALPN: &[u8] = b"doq";
#[derive(Clone)]
pub struct DoqBackendConfig {
pub server: SocketAddr,
pub server_name: ServerName<'static>,
pub tls_config: Arc<RustlsClientConfig>,
pub connect_timeout: Duration,
pub read_timeout: Duration,
}
impl DoqBackendConfig {
pub fn with_webpki_roots(
server: SocketAddr,
server_name: ServerName<'static>,
connect_timeout: Duration,
read_timeout: Duration,
) -> Self {
let mut root_store = rustls::RootCertStore::empty();
root_store.extend(webpki_roots::TLS_SERVER_ROOTS.iter().cloned());
let mut tls_config = RustlsClientConfig::builder()
.with_root_certificates(root_store)
.with_no_client_auth();
tls_config.alpn_protocols = vec![DOQ_ALPN.to_vec()];
Self {
server,
server_name,
tls_config: Arc::new(tls_config),
connect_timeout,
read_timeout,
}
}
}
pub struct DoqBackend {
config: DoqBackendConfig,
}
impl DoqBackend {
pub fn new(config: DoqBackendConfig) -> Self {
Self { config }
}
}
#[async_trait]
impl UpstreamBackend for DoqBackend {
async fn resolve(&self, query: &Message) -> Result<Message> {
let quic_client_config: QuicClientConfig =
self.config.tls_config.clone().try_into().map_err(
|err: quinn::crypto::rustls::NoInitialCipherSuite| Error::Tls(err.to_string()),
)?;
let client_config = ClientConfig::new(Arc::new(quic_client_config));
let bind_addr = unspecified_like(self.config.server);
let endpoint = Endpoint::client(bind_addr)
.map_err(|err| Error::Transport(format!("binding QUIC endpoint: {err}")))?;
let server_name = self.config.server_name.to_str();
let connecting = endpoint
.connect_with(client_config, self.config.server, server_name.as_ref())
.map_err(|err| Error::Transport(err.to_string()))?;
let connection = timeout(self.config.connect_timeout, connecting)
.await
.map_err(|_| Error::Timeout)?
.map_err(connection_error_to_lattice_error)?;
let (send, recv) = timeout(self.config.connect_timeout, connection.open_bi())
.await
.map_err(|_| Error::Timeout)?
.map_err(connection_error_to_lattice_error)?;
let mut stream = QuicStream { send, recv };
let response = framed_query(&mut stream, self.config.read_timeout, query).await;
let _ = stream.send.finish();
response
}
}
pub(crate) struct QuicStream {
pub(crate) send: SendStream,
pub(crate) recv: RecvStream,
}
impl AsyncRead for QuicStream {
fn poll_read(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<std::io::Result<()>> {
AsyncRead::poll_read(Pin::new(&mut self.recv), cx, buf)
}
}
impl AsyncWrite for QuicStream {
fn poll_write(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<std::io::Result<usize>> {
AsyncWrite::poll_write(Pin::new(&mut self.send), cx, buf)
}
fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
AsyncWrite::poll_flush(Pin::new(&mut self.send), cx)
}
fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
AsyncWrite::poll_shutdown(Pin::new(&mut self.send), cx)
}
}
fn unspecified_like(addr: SocketAddr) -> SocketAddr {
match addr {
SocketAddr::V4(_) => SocketAddr::new(IpAddr::V4(Ipv4Addr::UNSPECIFIED), 0),
SocketAddr::V6(_) => SocketAddr::new(IpAddr::V6(Ipv6Addr::UNSPECIFIED), 0),
}
}
fn connection_error_to_lattice_error(err: ConnectionError) -> Error {
let message = err.to_string();
if matches!(err, ConnectionError::TransportError(_))
&& message.contains("cryptographic handshake failed")
{
Error::Tls(message)
} else {
Error::Transport(message)
}
}
#[cfg(test)]
mod tests {
use super::*;
use dns_lattice_model::{Class, Header, Name, Opcode, Question, Rcode, RecordType};
use quinn::crypto::rustls::QuicServerConfig;
use quinn::{ServerConfig, TransportConfig};
use rcgen::{CertifiedKey, generate_simple_self_signed};
use rustls::pki_types::{CertificateDer, PrivateKeyDer};
use rustls::{RootCertStore, ServerConfig as RustlsServerConfig};
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, RustlsClientConfig, 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 mut rustls_server_config = RustlsServerConfig::builder()
.with_no_client_auth()
.with_single_cert(vec![cert_der.clone()], key_der)
.unwrap();
rustls_server_config.alpn_protocols = vec![DOQ_ALPN.to_vec()];
let quic_server_config: QuicServerConfig = rustls_server_config
.try_into()
.expect("valid TLS 1.3 initial cipher suite");
let mut server_config = ServerConfig::with_crypto(Arc::new(quic_server_config));
let mut transport = TransportConfig::default();
transport.max_idle_timeout(Some(Duration::from_secs(5).try_into().unwrap()));
server_config.transport_config(Arc::new(transport));
let mut roots = RootCertStore::empty();
roots.add(cert_der).unwrap();
let mut client_config = RustlsClientConfig::builder()
.with_root_certificates(roots)
.with_no_client_auth();
client_config.alpn_protocols = vec![DOQ_ALPN.to_vec()];
let server_name = ServerName::try_from("localhost").unwrap();
(server_config, client_config, server_name)
}
#[tokio::test]
async fn doq_backend_resolves_against_a_loopback_quic_server() {
let (server_config, client_config, server_name) = self_signed_fixture();
let endpoint = Endpoint::server(server_config, "127.0.0.1:0".parse().unwrap()).unwrap();
let addr = endpoint.local_addr().unwrap();
let responder = tokio::spawn(async move {
let incoming = endpoint.accept().await.unwrap();
let connection = incoming.await.unwrap();
let (mut send, mut recv) = connection.accept_bi().await.unwrap();
let mut len_buf = [0u8; 2];
recv.read_exact(&mut len_buf).await.unwrap();
let len = u16::from_be_bytes(len_buf) as usize;
let mut payload = vec![0u8; len];
recv.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);
send.write_all(&framed).await.unwrap();
let _ = send.finish();
let _ = send.stopped().await;
});
let backend = DoqBackend::new(DoqBackendConfig {
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("doq backend resolves");
assert!(answer.header.qr);
responder.await.unwrap();
}
#[tokio::test]
async fn doq_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 endpoint = Endpoint::server(server_config, "127.0.0.1:0".parse().unwrap()).unwrap();
let addr = endpoint.local_addr().unwrap();
let responder = tokio::spawn(async move {
if let Some(incoming) = endpoint.accept().await {
let _ = incoming.await;
}
});
let backend = DoqBackend::new(DoqBackendConfig {
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(_)),
"expected Error::Tls, got {err:?}"
);
responder.abort();
}
#[tokio::test]
async fn doq_backend_transport_error_on_connect_failure() {
let (_server_config, client_config, server_name) = self_signed_fixture();
let placeholder = tokio::net::UdpSocket::bind("127.0.0.1:0").await.unwrap();
let addr = placeholder.local_addr().unwrap();
drop(placeholder);
let backend = DoqBackend::new(DoqBackendConfig {
server: addr,
server_name,
tls_config: Arc::new(client_config),
connect_timeout: Duration::from_millis(200),
read_timeout: Duration::from_secs(2),
});
let err = backend
.resolve(&query_for("example.com"))
.await
.expect_err("connecting to a QUIC endpoint with no listener times out");
assert_eq!(err, Error::Timeout);
}
#[tokio::test]
async fn doq_backend_times_out_when_server_never_completes_handshake() {
let (server_config, client_config, server_name) = self_signed_fixture();
let _server_config = server_config;
let listener = tokio::net::UdpSocket::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let _keep_alive = listener;
let backend = DoqBackend::new(DoqBackendConfig {
server: addr,
server_name,
tls_config: Arc::new(client_config),
connect_timeout: Duration::from_millis(50),
read_timeout: Duration::from_secs(2),
});
let err = backend
.resolve(&query_for("example.com"))
.await
.expect_err("handshake does not complete within the timeout budget");
assert_eq!(err, Error::Timeout);
}
}