use std::net::SocketAddr;
use std::time::Duration;
use async_trait::async_trait;
use hickory_proto::op::{Message, update_message};
use hickory_proto::rr::rdata::TXT;
use hickory_proto::rr::rdata::tsig::TsigAlgorithm;
use hickory_proto::rr::{DNSClass, Name, RData, Record, RecordSet, RecordType, TSigner};
use hickory_proto::serialize::binary::{BinDecodable, BinEncodable};
use tokio::net::{TcpStream, UdpSocket};
use tracing::debug;
use crate::config::Rfc2136Config;
#[async_trait]
pub trait DnsUpdater: Send + Sync {
async fn upsert_txt(&self, name: &str, value: &str) -> Result<(), String>;
async fn delete_txt(&self, name: &str, value: &str) -> Result<(), String>;
}
pub struct Rfc2136Updater {
server: SocketAddr,
zone: Name,
signer: TSigner,
timeout: Duration,
ttl: u32,
}
const CHALLENGE_TTL: u32 = 60;
const UPDATE_TIMEOUT: Duration = Duration::from_secs(10);
impl Rfc2136Updater {
pub fn from_config(cfg: &Rfc2136Config) -> anyhow::Result<Self> {
use base64::prelude::*;
if cfg.server.is_empty() {
anyhow::bail!("signer.relay.dns01.rfc2136.server is not set");
}
let server: SocketAddr = std::net::ToSocketAddrs::to_socket_addrs(&cfg.server)
.map_err(|error| {
anyhow::anyhow!("rfc2136.server ({}) failed to resolve: {error}", cfg.server)
})?
.next()
.ok_or_else(|| {
anyhow::anyhow!("rfc2136.server ({}) resolved to no addresses", cfg.server)
})?;
let zone = Name::from_utf8(&cfg.zone).map_err(|error| {
anyhow::anyhow!("rfc2136.zone ({}) is not a DNS name: {error}", cfg.zone)
})?;
if cfg.tsig_key_name.is_empty() {
anyhow::bail!("signer.relay.dns01.rfc2136.tsig_key_name is not set");
}
let key_name = Name::from_utf8(&cfg.tsig_key_name)
.map_err(|error| anyhow::anyhow!("rfc2136.tsig_key_name is not a DNS name: {error}"))?;
let secret = BASE64_STANDARD
.decode(cfg.tsig_key_secret.trim())
.map_err(|error| anyhow::anyhow!("rfc2136.tsig_key_secret is not base64: {error}"))?;
if secret.is_empty() {
anyhow::bail!("signer.relay.dns01.rfc2136.tsig_key_secret is empty");
}
let algorithm = tsig_algorithm(&cfg.tsig_algorithm)?;
let signer = TSigner::new(secret, algorithm, key_name, 300)
.map_err(|error| anyhow::anyhow!("TSIG signer unusable: {error}"))?;
Ok(Self {
server,
zone,
signer,
timeout: UPDATE_TIMEOUT,
ttl: CHALLENGE_TTL,
})
}
fn txt_record(&self, name: &str, value: &str) -> Result<Record, String> {
let name =
Name::from_utf8(name).map_err(|error| format!("{name} is not a DNS name: {error}"))?;
let mut record = Record::from_rdata(
name,
self.ttl,
RData::TXT(TXT::new(vec![value.to_string()])),
);
record.dns_class = DNSClass::IN;
Ok(record)
}
async fn send(&self, mut message: Message) -> Result<(), String> {
use hickory_proto::op::ResponseCode;
let id = message.id;
message
.finalize(&self.signer, now_secs())
.map_err(|error| format!("signing the DNS update failed: {error}"))?;
let bytes = message
.to_bytes()
.map_err(|error| format!("encoding the DNS update failed: {error}"))?;
let response = tokio::time::timeout(self.timeout, self.exchange(&bytes))
.await
.map_err(|_| format!("DNS update to {} timed out", self.server))??;
let response = Message::from_bytes(&response)
.map_err(|error| format!("decoding the DNS response failed: {error}"))?;
if response.id != id {
return Err("DNS response id did not match the request".to_string());
}
match response.response_code {
ResponseCode::NoError => Ok(()),
other => Err(format!("DNS update refused: {other}")),
}
}
async fn exchange(&self, request: &[u8]) -> Result<Vec<u8>, String> {
let bind: SocketAddr = if self.server.is_ipv4() {
"0.0.0.0:0".parse().expect("a valid bind address")
} else {
"[::]:0".parse().expect("a valid bind address")
};
let socket = UdpSocket::bind(bind)
.await
.map_err(|error| format!("binding a UDP socket failed: {error}"))?;
socket
.send_to(request, self.server)
.await
.map_err(|error| format!("sending to {} failed: {error}", self.server))?;
let mut buffer = vec![0u8; 4096];
let read = socket
.recv(&mut buffer)
.await
.map_err(|error| format!("no answer from {}: {error}", self.server))?;
buffer.truncate(read);
if Message::from_bytes(&buffer)
.map(|message| message.truncation)
.unwrap_or(false)
{
debug!(
event = "signer_relay_dns_01_update_truncated",
outcome = "progress"
);
return self.exchange_tcp(request).await;
}
Ok(buffer)
}
async fn exchange_tcp(&self, request: &[u8]) -> Result<Vec<u8>, String> {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let mut stream = TcpStream::connect(self.server)
.await
.map_err(|error| format!("connecting to {} failed: {error}", self.server))?;
let length = u16::try_from(request.len())
.map_err(|_| "the DNS update is too large for TCP framing".to_string())?;
stream
.write_all(&length.to_be_bytes())
.await
.map_err(|error| format!("writing to {} failed: {error}", self.server))?;
stream
.write_all(request)
.await
.map_err(|error| format!("writing to {} failed: {error}", self.server))?;
let mut length = [0u8; 2];
stream
.read_exact(&mut length)
.await
.map_err(|error| format!("reading from {} failed: {error}", self.server))?;
let mut response = vec![0u8; u16::from_be_bytes(length) as usize];
stream
.read_exact(&mut response)
.await
.map_err(|error| format!("reading from {} failed: {error}", self.server))?;
Ok(response)
}
}
#[async_trait]
impl DnsUpdater for Rfc2136Updater {
async fn upsert_txt(&self, name: &str, value: &str) -> Result<(), String> {
let record = self.txt_record(name, value)?;
let mut rrset = RecordSet::new(record.name.clone(), RecordType::TXT, 0);
rrset.insert(record, 0);
let message = update_message::append(rrset, self.zone.clone(), false, true);
self.send(message).await
}
async fn delete_txt(&self, name: &str, value: &str) -> Result<(), String> {
let record = self.txt_record(name, value)?;
let message = update_message::delete_rrset(record, self.zone.clone(), true);
self.send(message).await
}
}
fn tsig_algorithm(name: &str) -> anyhow::Result<TsigAlgorithm> {
match name.trim().to_ascii_lowercase().as_str() {
"" | "hmac-sha256" => Ok(TsigAlgorithm::HmacSha256),
"hmac-sha384" => Ok(TsigAlgorithm::HmacSha384),
"hmac-sha512" => Ok(TsigAlgorithm::HmacSha512),
other => anyhow::bail!(
"unknown rfc2136.tsig_algorithm: {other} (supported: hmac-sha256, hmac-sha384, hmac-sha512)"
),
}
}
fn now_secs() -> u64 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_secs())
.unwrap_or(0)
}
#[cfg(test)]
mod tests {
use super::*;
fn config() -> Rfc2136Config {
use base64::prelude::*;
Rfc2136Config {
server: "127.0.0.1:53".to_string(),
zone: "example.org.".to_string(),
tsig_key_name: "acme-key.".to_string(),
tsig_key_secret: BASE64_STANDARD.encode(b"0123456789abcdef0123456789abcdef"),
tsig_algorithm: "hmac-sha256".to_string(),
}
}
#[test]
fn a_well_formed_config_builds() {
let updater = Rfc2136Updater::from_config(&config()).unwrap();
assert_eq!(updater.server.port(), 53);
assert_eq!(updater.zone.to_utf8(), "example.org.");
}
#[test]
fn every_malformed_field_is_a_startup_error() {
type Case = (&'static str, Box<dyn Fn(&mut Rfc2136Config)>);
let cases: Vec<Case> = vec![
(
"server",
Box::new(|c: &mut Rfc2136Config| c.server = String::new()),
),
(
"failed to resolve",
Box::new(|c: &mut Rfc2136Config| c.server = "not-an-address".to_string()),
),
(
"tsig_key_name",
Box::new(|c: &mut Rfc2136Config| c.tsig_key_name = String::new()),
),
(
"base64",
Box::new(|c: &mut Rfc2136Config| {
c.tsig_key_secret = "!!!not base64!!!".to_string()
}),
),
(
"empty",
Box::new(|c: &mut Rfc2136Config| c.tsig_key_secret = String::new()),
),
(
"tsig_algorithm",
Box::new(|c: &mut Rfc2136Config| c.tsig_algorithm = "hmac-md5".to_string()),
),
];
for (expected, mutate) in cases {
let mut cfg = config();
mutate(&mut cfg);
let error = Rfc2136Updater::from_config(&cfg)
.err()
.unwrap_or_else(|| panic!("{expected}: this configuration must not build"))
.to_string();
assert!(
error.contains(expected),
"expected {expected:?} in the error, got: {error}"
);
}
}
#[test]
fn the_algorithm_defaults_to_sha256() {
assert!(matches!(tsig_algorithm(""), Ok(TsigAlgorithm::HmacSha256)));
assert!(matches!(
tsig_algorithm("HMAC-SHA512"),
Ok(TsigAlgorithm::HmacSha512)
));
}
#[test]
fn a_record_carries_the_value_and_a_short_ttl() {
let updater = Rfc2136Updater::from_config(&config()).unwrap();
let record = updater
.txt_record("_acme-challenge.example.org.", "digest-value")
.unwrap();
assert_eq!(record.ttl, CHALLENGE_TTL);
assert_eq!(record.record_type(), RecordType::TXT);
match &record.data {
RData::TXT(txt) => {
assert_eq!(txt.to_string(), "digest-value");
}
other => panic!("expected a TXT record, got {other:?}"),
}
}
#[test]
fn a_malformed_record_name_is_rejected() {
let updater = Rfc2136Updater::from_config(&config()).unwrap();
assert!(updater.txt_record("not a dns name", "value").is_err());
}
mod stub {
use super::*;
use hickory_proto::op::{MessageType, OpCode, ResponseCode};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
pub(super) struct Server {
pub(super) addr: SocketAddr,
}
#[derive(Clone, Copy)]
pub(super) enum Udp {
Answer(ResponseCode),
Truncated,
WrongId,
Garbage,
}
pub(super) async fn spawn(udp: Udp) -> Server {
let (tcp, socket) = loop {
let tcp = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = tcp.local_addr().unwrap().port();
match UdpSocket::bind(("127.0.0.1", port)).await {
Ok(socket) => break (tcp, socket),
Err(_) => continue,
}
};
let addr = tcp.local_addr().unwrap();
tokio::spawn(async move {
let mut buffer = vec![0u8; 4096];
let (read, peer) = socket.recv_from(&mut buffer).await.unwrap();
let request = Message::from_bytes(&buffer[..read]).unwrap();
let bytes = match udp {
Udp::Garbage => b"definitely not DNS".to_vec(),
Udp::Answer(code) => reply(request.id, code, false),
Udp::WrongId => reply(request.id.wrapping_add(1), ResponseCode::NoError, false),
Udp::Truncated => reply(request.id, ResponseCode::NoError, true),
};
socket.send_to(&bytes, peer).await.unwrap();
});
tokio::spawn(async move {
let Ok((mut stream, _)) = tcp.accept().await else {
return;
};
let mut length = [0u8; 2];
if stream.read_exact(&mut length).await.is_err() {
return;
}
let mut request = vec![0u8; u16::from_be_bytes(length) as usize];
if stream.read_exact(&mut request).await.is_err() {
return;
}
let id = Message::from_bytes(&request).map(|m| m.id).unwrap_or(0);
let bytes = reply(id, ResponseCode::NoError, false);
let framed = u16::try_from(bytes.len()).unwrap().to_be_bytes();
let _ = stream.write_all(&framed).await;
let _ = stream.write_all(&bytes).await;
});
Server { addr }
}
fn reply(id: u16, code: ResponseCode, truncated: bool) -> Vec<u8> {
let mut message = Message::response(id, OpCode::Update);
message.metadata.message_type = MessageType::Response;
message.metadata.response_code = code;
message.metadata.truncation = truncated;
message.to_bytes().unwrap()
}
}
fn updater_for(addr: SocketAddr) -> Rfc2136Updater {
let mut cfg = config();
cfg.server = addr.to_string();
let mut updater = Rfc2136Updater::from_config(&cfg).unwrap();
updater.timeout = Duration::from_secs(5);
updater
}
#[tokio::test]
async fn an_accepted_update_succeeds() {
use hickory_proto::op::ResponseCode;
let server = stub::spawn(stub::Udp::Answer(ResponseCode::NoError)).await;
updater_for(server.addr)
.upsert_txt("_acme-challenge.example.org.", "digest-value")
.await
.expect("NOERROR is an accepted update");
}
#[tokio::test]
async fn a_retraction_reaches_the_server() {
use hickory_proto::op::ResponseCode;
let server = stub::spawn(stub::Udp::Answer(ResponseCode::NoError)).await;
updater_for(server.addr)
.delete_txt("_acme-challenge.example.org.", "digest-value")
.await
.expect("NOERROR is an accepted retraction");
}
#[tokio::test]
async fn a_refused_update_reports_the_response_code() {
use hickory_proto::op::ResponseCode;
let server = stub::spawn(stub::Udp::Answer(ResponseCode::Refused)).await;
let error = updater_for(server.addr)
.upsert_txt("_acme-challenge.example.org.", "digest-value")
.await
.expect_err("REFUSED is not an accepted update");
assert!(error.contains("DNS update refused"), "{error}");
assert!(error.contains("Refused"), "{error}");
}
#[tokio::test]
async fn a_truncated_answer_is_retried_over_tcp() {
let server = stub::spawn(stub::Udp::Truncated).await;
updater_for(server.addr)
.upsert_txt("_acme-challenge.example.org.", "digest-value")
.await
.expect("the TCP retry must carry the answer");
}
#[tokio::test]
async fn a_mismatched_response_id_is_rejected() {
let server = stub::spawn(stub::Udp::WrongId).await;
let error = updater_for(server.addr)
.upsert_txt("_acme-challenge.example.org.", "digest-value")
.await
.expect_err("a foreign id is not this update's answer");
assert!(error.contains("did not match"), "{error}");
}
#[tokio::test]
async fn an_undecodable_response_is_reported() {
let server = stub::spawn(stub::Udp::Garbage).await;
let error = updater_for(server.addr)
.upsert_txt("_acme-challenge.example.org.", "digest-value")
.await
.expect_err("garbage is not a DNS response");
assert!(error.contains("decoding the DNS response"), "{error}");
}
#[tokio::test]
async fn a_malformed_name_fails_before_the_socket() {
let updater = updater_for("127.0.0.1:1".parse().unwrap());
assert!(updater.upsert_txt("not a dns name", "v").await.is_err());
assert!(updater.delete_txt("not a dns name", "v").await.is_err());
}
#[tokio::test]
async fn an_unanswered_update_times_out() {
let mut cfg = config();
cfg.server = "127.0.0.1:1".to_string();
let mut updater = Rfc2136Updater::from_config(&cfg).unwrap();
updater.timeout = Duration::from_millis(100);
let error = updater
.upsert_txt("_acme-challenge.example.org.", "value")
.await
.unwrap_err();
assert!(
error.contains("timed out") || error.contains("failed"),
"{error}"
);
}
}