Skip to main content

dns_update_lite/providers/
rfc2136.rs

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