Skip to main content

dns_update/providers/
rfc2136.rs

1/*
2 * Copyright Stalwart Labs LLC See the COPYING
3 * file at the top-level directory of this distribution.
4 *
5 * Licensed under the Apache License, Version 2.0 <LICENSE-APACHE or
6 * https://www.apache.org/licenses/LICENSE-2.0> or the MIT license
7 * <LICENSE-MIT or https://opensource.org/licenses/MIT>, at your
8 * option. This file may not be copied, modified, or distributed
9 * except according to those terms.
10 */
11
12use crate::utils::split_caa_value;
13use crate::utils::strip_trailing_dot;
14use crate::utils::txt_chunks;
15use crate::{
16    CAARecord, DnsRecord, DnsRecordType, Error, IntoFqdn, MXRecord, SRVRecord, TLSARecord,
17    TlsaCertUsage, TlsaMatching, TlsaSelector,
18};
19use hickory_net::NetError;
20use hickory_net::client::{Client, ClientHandle};
21use hickory_net::runtime::TokioRuntimeProvider;
22use hickory_net::tcp::TcpClientStream;
23use hickory_net::udp::UdpClientStream;
24use hickory_net::xfer::DnsMultiplexer;
25use hickory_proto::ProtoError;
26use hickory_proto::dnssec::DnsSecError;
27use hickory_proto::op::ResponseCode;
28use hickory_proto::rr::rdata::caa::KeyValue;
29use hickory_proto::rr::rdata::tlsa::{CertUsage, Matching, Selector};
30use hickory_proto::rr::rdata::tsig::TsigAlgorithm;
31use hickory_proto::rr::rdata::{A, AAAA, CAA, CNAME, MX, NS, SRV, TLSA, TXT};
32use hickory_proto::rr::{DNSClass, Name, RData, Record, RecordSet, RecordType, TSigner};
33use std::net::{AddrParseError, SocketAddr};
34
35#[derive(Clone)]
36pub struct Rfc2136Provider {
37    addr: DnsAddress,
38    signer: Option<TSigner>,
39}
40
41#[derive(Clone, Copy, Debug, PartialEq, Eq)]
42pub enum DnsAddress {
43    Tcp(SocketAddr),
44    Udp(SocketAddr),
45}
46
47impl Rfc2136Provider {
48    pub(crate) fn new_tsig(
49        addr: impl TryInto<DnsAddress>,
50        key_name: impl AsRef<str>,
51        key: impl Into<Vec<u8>>,
52        algorithm: TsigAlgorithm,
53    ) -> crate::Result<Self> {
54        Ok(Rfc2136Provider {
55            addr: addr
56                .try_into()
57                .map_err(|_| Error::Parse("Invalid address".to_string()))?,
58            signer: Some(TSigner::new(
59                key.into(),
60                algorithm,
61                Name::from_ascii(key_name.as_ref())?,
62                60,
63            )?),
64        })
65    }
66
67    async fn connect(&self) -> crate::Result<Client<TokioRuntimeProvider>> {
68        self.connect_inner(self.signer.as_ref()).await
69    }
70
71    async fn connect_unsigned(&self) -> crate::Result<Client<TokioRuntimeProvider>> {
72        self.connect_inner(None).await
73    }
74
75    async fn connect_inner(
76        &self,
77        signer: Option<&TSigner>,
78    ) -> crate::Result<Client<TokioRuntimeProvider>> {
79        match &self.addr {
80            DnsAddress::Udp(addr) => {
81                let mut builder = UdpClientStream::builder(*addr, TokioRuntimeProvider::new());
82                if let Some(signer) = signer {
83                    builder = builder.with_signer(Some(signer.clone()));
84                }
85                let stream = builder.build();
86                let (client, bg) = Client::from_sender(stream);
87                tokio::spawn(bg);
88                Ok(client)
89            }
90            DnsAddress::Tcp(addr) => {
91                let (stream_future, sender) =
92                    TcpClientStream::new(*addr, None, None, TokioRuntimeProvider::new());
93                let stream = stream_future.await?;
94                let mut multiplexer = DnsMultiplexer::new(stream, sender);
95                if let Some(signer) = signer {
96                    multiplexer = multiplexer.with_signer(signer.clone());
97                }
98                let (client, bg) = Client::from_sender(multiplexer);
99                tokio::spawn(bg);
100                Ok(client)
101            }
102        }
103    }
104
105    pub(crate) async fn set_rrset(
106        &self,
107        name: impl IntoFqdn<'_>,
108        record_type: DnsRecordType,
109        ttl: u32,
110        records: Vec<DnsRecord>,
111        origin: impl IntoFqdn<'_>,
112    ) -> crate::Result<()> {
113        let owner = Name::from_str_relaxed(name.into_name().as_ref())?;
114        let zone = Name::from_str_relaxed(origin.into_fqdn().as_ref())?;
115        let rtype: RecordType = record_type.into();
116
117        let mut client = self.connect().await?;
118
119        let delete = Record::update0(owner.clone(), 0, rtype);
120        let result = client.delete_rrset(delete, zone.clone()).await?;
121        if result.response_code != ResponseCode::NoError {
122            return Err(Error::Response(result.response_code.to_string()));
123        }
124
125        if records.is_empty() {
126            return Ok(());
127        }
128
129        let rrset = build_rrset(owner, rtype, ttl, records)?;
130        let result = client.append(rrset, zone, false).await?;
131        if result.response_code != ResponseCode::NoError {
132            return Err(Error::Response(result.response_code.to_string()));
133        }
134        Ok(())
135    }
136
137    pub(crate) async fn add_to_rrset(
138        &self,
139        name: impl IntoFqdn<'_>,
140        record_type: DnsRecordType,
141        ttl: u32,
142        records: Vec<DnsRecord>,
143        origin: impl IntoFqdn<'_>,
144    ) -> crate::Result<()> {
145        if records.is_empty() {
146            return Ok(());
147        }
148        let owner = Name::from_str_relaxed(name.into_name().as_ref())?;
149        let zone = Name::from_str_relaxed(origin.into_fqdn().as_ref())?;
150        let rtype: RecordType = record_type.into();
151        let rrset = build_rrset(owner, rtype, ttl, records)?;
152
153        let mut client = self.connect().await?;
154        let result = client.append(rrset, zone, false).await?;
155        if result.response_code != ResponseCode::NoError {
156            return Err(Error::Response(result.response_code.to_string()));
157        }
158        Ok(())
159    }
160
161    pub(crate) async fn remove_from_rrset(
162        &self,
163        name: impl IntoFqdn<'_>,
164        record_type: DnsRecordType,
165        records: Vec<DnsRecord>,
166        origin: impl IntoFqdn<'_>,
167    ) -> crate::Result<()> {
168        if records.is_empty() {
169            return Ok(());
170        }
171        let owner = Name::from_str_relaxed(name.into_name().as_ref())?;
172        let zone = Name::from_str_relaxed(origin.into_fqdn().as_ref())?;
173        let rtype: RecordType = record_type.into();
174        let rrset = build_rrset(owner, rtype, 0, records)?;
175
176        let mut client = self.connect().await?;
177        let result = client.delete_by_rdata(rrset, zone).await?;
178        if result.response_code != ResponseCode::NoError {
179            return Err(Error::Response(result.response_code.to_string()));
180        }
181        Ok(())
182    }
183
184    pub(crate) async fn list_rrset(
185        &self,
186        name: impl IntoFqdn<'_>,
187        record_type: DnsRecordType,
188        _origin: impl IntoFqdn<'_>,
189    ) -> crate::Result<Vec<DnsRecord>> {
190        let owner = Name::from_str_relaxed(name.into_fqdn().as_ref())?;
191        let rtype: RecordType = record_type.into();
192
193        let mut client = self.connect_unsigned().await?;
194        let response = client.query(owner.clone(), DNSClass::IN, rtype).await?;
195        if response.response_code != ResponseCode::NoError
196            && response.response_code != ResponseCode::NXDomain
197        {
198            return Err(Error::Response(response.response_code.to_string()));
199        }
200
201        let mut out = Vec::new();
202        for record in response.answers.iter() {
203            if record.record_type() != rtype || record.name != owner {
204                continue;
205            }
206            out.push(rdata_to_dns_record(&record.data)?);
207        }
208        Ok(out)
209    }
210}
211
212fn rdata_to_dns_record(data: &RData) -> crate::Result<DnsRecord> {
213    Ok(match data {
214        RData::A(a) => DnsRecord::A(a.0),
215        RData::AAAA(aaaa) => DnsRecord::AAAA(aaaa.0),
216        RData::CNAME(cname) => DnsRecord::CNAME(strip_trailing_dot(&cname.0.to_utf8()).to_string()),
217        RData::NS(ns) => DnsRecord::NS(strip_trailing_dot(&ns.0.to_utf8()).to_string()),
218        RData::MX(mx) => DnsRecord::MX(MXRecord {
219            priority: mx.preference,
220            exchange: strip_trailing_dot(&mx.exchange.to_utf8()).to_string(),
221        }),
222        RData::TXT(txt) => {
223            let combined: String = txt
224                .txt_data
225                .iter()
226                .map(|chunk| String::from_utf8_lossy(chunk).into_owned())
227                .collect();
228            DnsRecord::TXT(combined)
229        }
230        RData::SRV(srv) => DnsRecord::SRV(SRVRecord {
231            priority: srv.priority,
232            weight: srv.weight,
233            port: srv.port,
234            target: strip_trailing_dot(&srv.target.to_utf8()).to_string(),
235        }),
236        RData::TLSA(tlsa) => DnsRecord::TLSA(TLSARecord {
237            cert_usage: tlsa_cert_usage_from(tlsa.cert_usage)?,
238            selector: tlsa_selector_from(tlsa.selector)?,
239            matching: tlsa_matching_from(tlsa.matching)?,
240            cert_data: tlsa.cert_data.clone(),
241        }),
242        RData::CAA(caa) => DnsRecord::CAA(caa_to_record(caa)?),
243        other => {
244            return Err(Error::Unsupported(format!(
245                "Unsupported RData type for list_rrset: {}",
246                other.record_type()
247            )));
248        }
249    })
250}
251
252fn caa_to_record(caa: &CAA) -> crate::Result<CAARecord> {
253    let issuer_critical = caa.issuer_critical;
254    let value_text = String::from_utf8_lossy(&caa.value).into_owned();
255    match caa.tag.as_str() {
256        "issue" => {
257            let (name, options) = split_caa_value(&value_text);
258            Ok(CAARecord::Issue {
259                issuer_critical,
260                name,
261                options,
262            })
263        }
264        "issuewild" => {
265            let (name, options) = split_caa_value(&value_text);
266            Ok(CAARecord::IssueWild {
267                issuer_critical,
268                name,
269                options,
270            })
271        }
272        "iodef" => Ok(CAARecord::Iodef {
273            issuer_critical,
274            url: value_text,
275        }),
276        other => Err(Error::Unsupported(format!(
277            "Unsupported CAA tag for list_rrset: {other}"
278        ))),
279    }
280}
281
282fn tlsa_cert_usage_from(usage: CertUsage) -> crate::Result<TlsaCertUsage> {
283    Ok(match usage {
284        CertUsage::PkixTa => TlsaCertUsage::PkixTa,
285        CertUsage::PkixEe => TlsaCertUsage::PkixEe,
286        CertUsage::DaneTa => TlsaCertUsage::DaneTa,
287        CertUsage::DaneEe => TlsaCertUsage::DaneEe,
288        CertUsage::Private => TlsaCertUsage::Private,
289        other => return Err(Error::Api(format!("Unknown TLSA cert usage: {other:?}"))),
290    })
291}
292
293fn tlsa_selector_from(sel: Selector) -> crate::Result<TlsaSelector> {
294    Ok(match sel {
295        Selector::Full => TlsaSelector::Full,
296        Selector::Spki => TlsaSelector::Spki,
297        Selector::Private => TlsaSelector::Private,
298        other => return Err(Error::Api(format!("Unknown TLSA selector: {other:?}"))),
299    })
300}
301
302fn tlsa_matching_from(m: Matching) -> crate::Result<TlsaMatching> {
303    Ok(match m {
304        Matching::Raw => TlsaMatching::Raw,
305        Matching::Sha256 => TlsaMatching::Sha256,
306        Matching::Sha512 => TlsaMatching::Sha512,
307        Matching::Private => TlsaMatching::Private,
308        other => return Err(Error::Api(format!("Unknown TLSA matching: {other:?}"))),
309    })
310}
311
312fn build_rrset(
313    name: Name,
314    rtype: RecordType,
315    ttl: u32,
316    records: Vec<DnsRecord>,
317) -> crate::Result<RecordSet> {
318    let mut rrset = RecordSet::with_ttl(name, rtype, ttl);
319    for record in records {
320        let (record_type, rdata) = convert_record(record)?;
321        if record_type != rtype {
322            return Err(Error::Api(format!(
323                "RRSet record type mismatch: expected {rtype}, got {record_type}"
324            )));
325        }
326        rrset.add_rdata(rdata);
327    }
328    Ok(rrset)
329}
330
331impl From<DnsRecordType> for RecordType {
332    fn from(record_type: DnsRecordType) -> Self {
333        match record_type {
334            DnsRecordType::A => RecordType::A,
335            DnsRecordType::AAAA => RecordType::AAAA,
336            DnsRecordType::CNAME => RecordType::CNAME,
337            DnsRecordType::NS => RecordType::NS,
338            DnsRecordType::MX => RecordType::MX,
339            DnsRecordType::TXT => RecordType::TXT,
340            DnsRecordType::SRV => RecordType::SRV,
341            DnsRecordType::TLSA => RecordType::TLSA,
342            DnsRecordType::CAA => RecordType::CAA,
343        }
344    }
345}
346
347fn convert_record(record: DnsRecord) -> crate::Result<(RecordType, RData)> {
348    Ok(match record {
349        DnsRecord::A(content) => (RecordType::A, RData::A(A::from(content))),
350        DnsRecord::AAAA(content) => (RecordType::AAAA, RData::AAAA(AAAA::from(content))),
351        DnsRecord::CNAME(content) => (
352            RecordType::CNAME,
353            RData::CNAME(CNAME(Name::from_str_relaxed(content)?)),
354        ),
355        DnsRecord::NS(content) => (
356            RecordType::NS,
357            RData::NS(NS(Name::from_str_relaxed(content)?)),
358        ),
359        DnsRecord::MX(content) => (
360            RecordType::MX,
361            RData::MX(MX::new(
362                content.priority,
363                Name::from_str_relaxed(content.exchange)?,
364            )),
365        ),
366        DnsRecord::TXT(content) => (RecordType::TXT, RData::TXT(TXT::new(txt_chunks(content)))),
367        DnsRecord::SRV(content) => (
368            RecordType::SRV,
369            RData::SRV(SRV::new(
370                content.priority,
371                content.weight,
372                content.port,
373                Name::from_str_relaxed(content.target)?,
374            )),
375        ),
376        DnsRecord::TLSA(content) => (
377            RecordType::TLSA,
378            RData::TLSA(TLSA::new(
379                content.cert_usage.into(),
380                content.selector.into(),
381                content.matching.into(),
382                content.cert_data,
383            )),
384        ),
385        DnsRecord::CAA(caa) => (
386            RecordType::CAA,
387            RData::CAA(match caa {
388                CAARecord::Issue {
389                    issuer_critical,
390                    name,
391                    options,
392                } => CAA::new_issue(
393                    issuer_critical,
394                    name.map(Name::from_str_relaxed).transpose()?,
395                    options
396                        .into_iter()
397                        .map(|kv| KeyValue::new(kv.key, kv.value))
398                        .collect(),
399                ),
400                CAARecord::IssueWild {
401                    issuer_critical,
402                    name,
403                    options,
404                } => CAA::new_issuewild(
405                    issuer_critical,
406                    name.map(Name::from_str_relaxed).transpose()?,
407                    options
408                        .into_iter()
409                        .map(|kv| KeyValue::new(kv.key, kv.value))
410                        .collect(),
411                ),
412                CAARecord::Iodef {
413                    issuer_critical,
414                    url,
415                } => CAA::new_iodef(
416                    issuer_critical,
417                    url.parse()
418                        .map_err(|_| Error::Parse("Invalid URL in CAA record".to_string()))?,
419                ),
420            }),
421        ),
422    })
423}
424
425impl From<TlsaCertUsage> for CertUsage {
426    fn from(usage: TlsaCertUsage) -> Self {
427        match usage {
428            TlsaCertUsage::PkixTa => CertUsage::PkixTa,
429            TlsaCertUsage::PkixEe => CertUsage::PkixEe,
430            TlsaCertUsage::DaneTa => CertUsage::DaneTa,
431            TlsaCertUsage::DaneEe => CertUsage::DaneEe,
432            TlsaCertUsage::Private => CertUsage::Private,
433        }
434    }
435}
436
437impl From<TlsaMatching> for Matching {
438    fn from(matching: TlsaMatching) -> Self {
439        match matching {
440            TlsaMatching::Raw => Matching::Raw,
441            TlsaMatching::Sha256 => Matching::Sha256,
442            TlsaMatching::Sha512 => Matching::Sha512,
443            TlsaMatching::Private => Matching::Private,
444        }
445    }
446}
447
448impl From<TlsaSelector> for Selector {
449    fn from(selector: TlsaSelector) -> Self {
450        match selector {
451            TlsaSelector::Full => Selector::Full,
452            TlsaSelector::Spki => Selector::Spki,
453            TlsaSelector::Private => Selector::Private,
454        }
455    }
456}
457
458impl TryFrom<&str> for DnsAddress {
459    type Error = ();
460
461    fn try_from(url: &str) -> Result<Self, Self::Error> {
462        let (host, is_tcp) = if let Some(host) = url.strip_prefix("udp://") {
463            (host, false)
464        } else if let Some(host) = url.strip_prefix("tcp://") {
465            (host, true)
466        } else {
467            (url, false)
468        };
469        let (host, port) = if let Some(host) = host.strip_prefix('[') {
470            let (host, maybe_port) = host.rsplit_once(']').ok_or(())?;
471
472            (
473                host,
474                maybe_port
475                    .rsplit_once(':')
476                    .map(|(_, port)| port)
477                    .unwrap_or("53"),
478            )
479        } else if let Some((host, port)) = host.rsplit_once(':') {
480            (host, port)
481        } else {
482            (host, "53")
483        };
484
485        let addr = SocketAddr::new(host.parse().map_err(|_| ())?, port.parse().map_err(|_| ())?);
486
487        if is_tcp {
488            Ok(DnsAddress::Tcp(addr))
489        } else {
490            Ok(DnsAddress::Udp(addr))
491        }
492    }
493}
494
495impl TryFrom<&String> for DnsAddress {
496    type Error = ();
497
498    fn try_from(url: &String) -> Result<Self, Self::Error> {
499        DnsAddress::try_from(url.as_str())
500    }
501}
502
503impl TryFrom<String> for DnsAddress {
504    type Error = ();
505
506    fn try_from(url: String) -> Result<Self, Self::Error> {
507        DnsAddress::try_from(url.as_str())
508    }
509}
510
511impl From<crate::TsigAlgorithm> for TsigAlgorithm {
512    fn from(alg: crate::TsigAlgorithm) -> Self {
513        match alg {
514            crate::TsigAlgorithm::HmacMd5 => TsigAlgorithm::HmacMd5,
515            crate::TsigAlgorithm::Gss => TsigAlgorithm::Gss,
516            crate::TsigAlgorithm::HmacSha1 => TsigAlgorithm::HmacSha1,
517            crate::TsigAlgorithm::HmacSha224 => TsigAlgorithm::HmacSha224,
518            crate::TsigAlgorithm::HmacSha256 => TsigAlgorithm::HmacSha256,
519            crate::TsigAlgorithm::HmacSha256_128 => TsigAlgorithm::HmacSha256_128,
520            crate::TsigAlgorithm::HmacSha384 => TsigAlgorithm::HmacSha384,
521            crate::TsigAlgorithm::HmacSha384_192 => TsigAlgorithm::HmacSha384_192,
522            crate::TsigAlgorithm::HmacSha512 => TsigAlgorithm::HmacSha512,
523            crate::TsigAlgorithm::HmacSha512_256 => TsigAlgorithm::HmacSha512_256,
524        }
525    }
526}
527
528impl From<ProtoError> for Error {
529    fn from(e: ProtoError) -> Self {
530        Error::Protocol(e.to_string())
531    }
532}
533
534impl From<AddrParseError> for Error {
535    fn from(e: AddrParseError) -> Self {
536        Error::Parse(e.to_string())
537    }
538}
539
540impl From<NetError> for Error {
541    fn from(e: NetError) -> Self {
542        Error::Client(e.to_string())
543    }
544}
545
546impl From<DnsSecError> for Error {
547    fn from(e: DnsSecError) -> Self {
548        Error::Protocol(e.to_string())
549    }
550}