Skip to main content

mdns_sd/
dns_parser.rs

1//! DNS parsing utility.
2//!
3//! [DnsIncoming] is the logic representation of an incoming DNS packet.
4//! [DnsOutgoing] is the logic representation of an outgoing DNS message of one or more packets.
5//! [DnsOutPacket] is the encoded one packet for [DnsOutgoing].
6
7#[cfg(feature = "logging")]
8use crate::log::{debug, trace};
9
10use crate::error::{e_fmt, Error, Result};
11use crate::service_info::{decode_txt, is_unicast_link_local, DnsRegistry, MyIntf, ServiceInfo};
12
13use if_addrs::Interface;
14
15#[cfg(feature = "serde")]
16use serde::{Deserialize, Serialize};
17
18use std::{
19    any::Any,
20    cmp,
21    collections::HashMap,
22    convert::TryInto,
23    fmt,
24    hash::Hash,
25    net::{IpAddr, Ipv4Addr, Ipv6Addr},
26    str,
27    time::{Duration, Instant},
28};
29
30/// Represents a network interface identifier defined by the OS.
31#[derive(Clone, Debug, Eq, Hash, PartialEq, Default)]
32#[cfg_attr(feature = "serde", derive(Deserialize, Serialize))]
33pub struct InterfaceId {
34    /// Interface name, e.g. "en0", "wlan0", etc.
35    pub name: String,
36
37    /// Interface index assigned by the OS, e.g. 1, 2, etc.
38    pub index: u32,
39}
40
41impl InterfaceId {
42    /// Returns all IP addresses associated with this interface by querying the OS.
43    pub fn get_addrs(&self) -> Vec<IpAddr> {
44        if_addrs::get_if_addrs()
45            .unwrap_or_default()
46            .into_iter()
47            .filter(|iface| iface.index == Some(self.index))
48            .map(|iface| iface.ip())
49            .collect()
50    }
51}
52
53impl fmt::Display for InterfaceId {
54    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
55        write!(f, "{}('{}')", self.index, self.name)
56    }
57}
58
59impl From<&Interface> for InterfaceId {
60    fn from(interface: &Interface) -> Self {
61        InterfaceId {
62            name: interface.name.clone(),
63            index: interface.index.unwrap_or_default(),
64        }
65    }
66}
67
68/// An IPv4 address with interface identifiers indicating which interfaces discovered it.
69#[derive(Debug, Clone, Eq, PartialEq, Hash)]
70#[cfg_attr(feature = "serde", derive(Deserialize, Serialize))]
71pub struct ScopedIpV4 {
72    addr: Ipv4Addr,
73    /// The interfaces this address was discovered on.
74    interface_ids: Vec<InterfaceId>,
75}
76
77impl ScopedIpV4 {
78    /// Creates a new `ScopedIpV4` with a single interface identifier.
79    pub fn new(addr: Ipv4Addr, interface_id: InterfaceId) -> Self {
80        Self {
81            addr,
82            interface_ids: vec![interface_id],
83        }
84    }
85
86    /// Returns the IPv4 address.
87    pub const fn addr(&self) -> &Ipv4Addr {
88        &self.addr
89    }
90
91    /// Returns the interfaces this address was discovered on.
92    pub fn interface_ids(&self) -> &[InterfaceId] {
93        &self.interface_ids
94    }
95
96    /// Adds an interface identifier if not already present.
97    pub(crate) fn add_interface_id(&mut self, id: InterfaceId) {
98        if !self.interface_ids.contains(&id) {
99            self.interface_ids.push(id);
100        }
101    }
102}
103
104/// An IPv6 address with scope_id (interface identifier).
105#[derive(Debug, Clone, Eq, PartialEq, Hash)]
106#[cfg_attr(feature = "serde", derive(Deserialize, Serialize))]
107pub struct ScopedIpV6 {
108    addr: Ipv6Addr,
109    scope_id: InterfaceId,
110}
111
112impl ScopedIpV6 {
113    /// Returns the IPv6 address.
114    pub const fn addr(&self) -> &Ipv6Addr {
115        &self.addr
116    }
117
118    /// Returns the scope_id for this IPv6 address.
119    pub const fn scope_id(&self) -> &InterfaceId {
120        &self.scope_id
121    }
122}
123
124/// An IP address, either IPv4 or IPv6, that supports scope_id for IPv6.
125#[derive(Debug, Clone, Eq, PartialEq, Hash)]
126#[cfg_attr(feature = "serde", derive(Deserialize, Serialize))]
127#[non_exhaustive]
128pub enum ScopedIp {
129    V4(ScopedIpV4),
130    V6(ScopedIpV6),
131}
132
133impl ScopedIp {
134    pub const fn to_ip_addr(&self) -> IpAddr {
135        match self {
136            ScopedIp::V4(v4) => IpAddr::V4(v4.addr),
137            ScopedIp::V6(v6) => IpAddr::V6(v6.addr),
138        }
139    }
140
141    pub const fn is_ipv4(&self) -> bool {
142        matches!(self, ScopedIp::V4(_))
143    }
144
145    pub const fn is_ipv6(&self) -> bool {
146        matches!(self, ScopedIp::V6(_))
147    }
148
149    pub const fn is_loopback(&self) -> bool {
150        match self {
151            ScopedIp::V4(v4) => v4.addr.is_loopback(),
152            ScopedIp::V6(v6) => v6.addr.is_loopback(),
153        }
154    }
155}
156
157impl From<IpAddr> for ScopedIp {
158    fn from(ip: IpAddr) -> Self {
159        match ip {
160            IpAddr::V4(v4) => ScopedIp::V4(ScopedIpV4 {
161                addr: v4,
162                interface_ids: vec![],
163            }),
164            IpAddr::V6(v6) => ScopedIp::V6(ScopedIpV6 {
165                addr: v6,
166                scope_id: InterfaceId::default(),
167            }),
168        }
169    }
170}
171
172impl From<&Interface> for ScopedIp {
173    fn from(interface: &Interface) -> Self {
174        match interface.ip() {
175            IpAddr::V4(v4) => ScopedIp::V4(ScopedIpV4 {
176                addr: v4,
177                interface_ids: vec![InterfaceId::from(interface)],
178            }),
179            IpAddr::V6(v6) => ScopedIp::V6(ScopedIpV6 {
180                addr: v6,
181                scope_id: InterfaceId::from(interface),
182            }),
183        }
184    }
185}
186
187impl fmt::Display for ScopedIp {
188    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
189        match self {
190            ScopedIp::V4(v4) => write!(f, "{}", v4.addr),
191            ScopedIp::V6(v6) => {
192                if v6.scope_id.index != 0 && is_unicast_link_local(&v6.addr) {
193                    #[cfg(windows)]
194                    {
195                        write!(f, "{}%{}", v6.addr, v6.scope_id.index)
196                    }
197                    #[cfg(not(windows))]
198                    {
199                        write!(f, "{}%{}", v6.addr, v6.scope_id.name)
200                    }
201                } else {
202                    write!(f, "{}", v6.addr)
203                }
204            }
205        }
206    }
207}
208
209/// DNS resource record types, stored as `u16`. Can do `as u16` when needed.
210///
211/// See [RFC 1035 section 3.2.2](https://datatracker.ietf.org/doc/html/rfc1035#section-3.2.2)
212#[derive(Debug, PartialEq, Eq, Clone, Copy, PartialOrd, Ord)]
213#[non_exhaustive]
214#[repr(u16)]
215pub enum RRType {
216    /// DNS record type for IPv4 address
217    A = 1,
218
219    /// DNS record type for Canonical Name
220    CNAME = 5,
221
222    /// DNS record type for Pointer
223    PTR = 12,
224
225    /// DNS record type for Host Info
226    HINFO = 13,
227
228    /// DNS record type for Text (properties)
229    TXT = 16,
230
231    /// DNS record type for IPv6 address
232    AAAA = 28,
233
234    /// DNS record type for Service
235    SRV = 33,
236
237    /// DNS record type for Negative Responses
238    NSEC = 47,
239
240    /// DNS service binding record
241    SVCB = 64,
242
243    /// HTTPS service binding record
244    HTTPS = 65,
245
246    /// DNS record type for any records (wildcard)
247    ANY = 255,
248}
249
250impl RRType {
251    /// Converts `u16` into `RRType` if possible.
252    pub const fn from_u16(value: u16) -> Option<Self> {
253        match value {
254            1 => Some(RRType::A),
255            5 => Some(RRType::CNAME),
256            12 => Some(RRType::PTR),
257            13 => Some(RRType::HINFO),
258            16 => Some(RRType::TXT),
259            28 => Some(RRType::AAAA),
260            33 => Some(RRType::SRV),
261            47 => Some(RRType::NSEC),
262            64 => Some(RRType::SVCB),
263            65 => Some(RRType::HTTPS),
264            255 => Some(RRType::ANY),
265            _ => None,
266        }
267    }
268}
269
270impl fmt::Display for RRType {
271    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
272        match self {
273            RRType::A => write!(f, "TYPE_A"),
274            RRType::CNAME => write!(f, "TYPE_CNAME"),
275            RRType::PTR => write!(f, "TYPE_PTR"),
276            RRType::HINFO => write!(f, "TYPE_HINFO"),
277            RRType::TXT => write!(f, "TYPE_TXT"),
278            RRType::AAAA => write!(f, "TYPE_AAAA"),
279            RRType::SRV => write!(f, "TYPE_SRV"),
280            RRType::NSEC => write!(f, "TYPE_NSEC"),
281            RRType::SVCB => write!(f, "TYPE_SVCB"),
282            RRType::HTTPS => write!(f, "TYPE_HTTPS"),
283            RRType::ANY => write!(f, "TYPE_ANY"),
284        }
285    }
286}
287
288/// The class value for the Internet.
289pub const CLASS_IN: u16 = 1;
290pub const CLASS_MASK: u16 = 0x7FFF;
291
292/// Cache-flush bit: the most significant bit of the rrclass field of the resource record.  
293pub const CLASS_CACHE_FLUSH: u16 = 0x8000;
294
295/// RFC 6762 ยง6.7: The resource record TTL given in a legacy unicast response SHOULD NOT
296/// be greater than ten seconds.
297pub const LEGACY_UNICAST_MAX_TTL: u32 = 10;
298
299/// Absolute max size of UDP datagram payload for an mDNS packet over IPv4.
300///
301/// RFC 6762 section 17:
302/// "Even when fragmentation is used, a Multicast DNS packet, including IP and UDP
303/// headers, MUST NOT exceed 9000 bytes."
304///
305/// It is calculated as: 9000 bytes - IPv4 header 20 bytes - UDP header 8 bytes.
306pub(crate) const MAX_PKT_ABSOLUTE_IPV4: usize = 8972;
307
308/// Absolute max size of UDP datagram payload for an mDNS packet over IPv6.
309///
310/// Same 9000-byte ceiling as [`MAX_PKT_ABSOLUTE_IPV4`], less the bigger IPv6 header:
311/// 9000 bytes - IPv6 header 40 bytes - UDP header 8 bytes.
312pub(crate) const MAX_PKT_ABSOLUTE_IPV6: usize = 8952;
313
314/// Absolute max size of an mDNS packet for the given IP version.
315pub(crate) const fn max_pkt_absolute(is_ipv4: bool) -> usize {
316    if is_ipv4 {
317        MAX_PKT_ABSOLUTE_IPV4
318    } else {
319        MAX_PKT_ABSOLUTE_IPV6
320    }
321}
322
323/// Default max size of a generated (i.e. outgoing) packet.
324///
325/// Calculated as: 1500 bytes Ethernet MTU - IPv6 header 40 bytes - UDP header 8 bytes.
326/// It is safe on both IPv4 and IPv6, at the cost of 20 unused bytes for IPv4.
327///
328/// The idea is to keep generated packets unfragmented at IP layer. See RFC 6762 section 17.
329pub const MAX_PKT_DEFAULT: usize = 1452;
330
331const MSG_HEADER_LEN: usize = 12;
332
333/// Max size of a single DNS label, in bytes.
334///
335/// Reference: [RFC1035 section 2.3.4](https://datatracker.ietf.org/doc/html/rfc1035#section-2.3.4)
336const MAX_LABEL_BYTES: usize = 63;
337
338/// Max size of a whole domain name, in bytes.
339///
340/// Reference: [RFC1035 section 2.3.4](https://datatracker.ietf.org/doc/html/rfc1035#section-2.3.4)
341const MAX_NAME_BYTES: usize = 255;
342
343/// Why a question or a record could not be written into a packet.
344///
345/// In either case nothing is left behind in the packet: the caller rolls back
346/// whatever was written and skips the item.
347#[derive(Debug, PartialEq, Eq)]
348pub enum WriteError {
349    /// A label in a name is longer than [`MAX_LABEL_BYTES`].
350    NameTooLong,
351
352    /// The packet would exceed its max size with this record.
353    PacketFull,
354}
355
356/// `crate::error::Result` shadows the std alias here, hence the full path.
357type WriteResult = core::result::Result<(), WriteError>;
358
359// Definitions for DNS message header "flags" field
360//
361// The "flags" field is 16-bit long, in this format:
362// (RFC 1035 section 4.1.1)
363//
364//   0  1  2  3  4  5  6  7  8  9  0  1  2  3  4  5
365// |QR|   Opcode  |AA|TC|RD|RA|   Z    |   RCODE   |
366//
367pub const FLAGS_QR_MASK: u16 = 0x8000; // mask for query/response bit
368
369/// Flag bit to indicate a query
370pub const FLAGS_QR_QUERY: u16 = 0x0000;
371
372/// Flag bit to indicate a response
373pub const FLAGS_QR_RESPONSE: u16 = 0x8000;
374
375/// Flag bit for Authoritative Answer
376pub const FLAGS_AA: u16 = 0x0400;
377
378/// mask for TC(Truncated) bit
379///
380/// 2024-08-10: currently this flag is only supported on the querier side,
381///             not supported on the responder side. I.e. the responder only
382///             handles the first packet and ignore this bit. Since the
383///             additional packets have 0 questions, the processing of them
384///             is no-op.
385///             In practice, this means the responder supports Known-Answer
386///             only with single packet, not multi-packet. The querier supports
387///             both single packet and multi-packet.
388pub const FLAGS_TC: u16 = 0x0200;
389
390/// A convenience type alias for DNS record trait objects.
391pub type DnsRecordBox = Box<dyn DnsRecordExt>;
392
393impl Clone for DnsRecordBox {
394    fn clone(&self) -> Self {
395        self.clone_box()
396    }
397}
398
399const U16_SIZE: usize = 2;
400
401/// Returns `RRType` for a given IP address.
402#[inline]
403pub const fn ip_address_rr_type(address: &IpAddr) -> RRType {
404    match address {
405        IpAddr::V4(_) => RRType::A,
406        IpAddr::V6(_) => RRType::AAAA,
407    }
408}
409
410#[derive(Eq, PartialEq, Debug, Clone)]
411pub struct DnsEntry {
412    pub(crate) name: String, // always lower case.
413    pub(crate) ty: RRType,
414    class: u16,
415    cache_flush: bool,
416}
417
418impl DnsEntry {
419    const fn new(name: String, ty: RRType, class: u16) -> Self {
420        Self {
421            name,
422            ty,
423            class: class & CLASS_MASK,
424            cache_flush: (class & CLASS_CACHE_FLUSH) != 0,
425        }
426    }
427}
428
429/// Common methods for all DNS entries:  questions and resource records.
430pub trait DnsEntryExt: fmt::Debug {
431    fn entry_name(&self) -> &str;
432
433    fn entry_type(&self) -> RRType;
434}
435
436/// A DNS question entry
437#[derive(Debug)]
438pub struct DnsQuestion {
439    pub(crate) entry: DnsEntry,
440}
441
442impl DnsEntryExt for DnsQuestion {
443    fn entry_name(&self) -> &str {
444        &self.entry.name
445    }
446
447    fn entry_type(&self) -> RRType {
448        self.entry.ty
449    }
450}
451
452/// A DNS Resource Record - like a DNS entry, but has a TTL.
453/// RFC: https://www.rfc-editor.org/rfc/rfc1035#section-3.2.1
454///      https://www.rfc-editor.org/rfc/rfc1035#section-4.1.3
455#[derive(Debug, Clone)]
456pub struct DnsRecord {
457    pub(crate) entry: DnsEntry,
458    ttl: u32, // in seconds, 0 means this record should not be cached
459    /// When this record was created (received or registered).
460    created: Instant,
461    /// When this record expires.
462    expires: Instant,
463
464    /// Support re-query an instance before its PTR record expires.
465    /// See https://datatracker.ietf.org/doc/html/rfc6762#section-5.2
466    refresh: Instant,
467
468    /// If conflict resolution decides to change the name, this is the new one.
469    new_name: Option<String>,
470}
471
472impl DnsRecord {
473    fn new(name: &str, ty: RRType, class: u16, ttl: u32) -> Self {
474        let created = Instant::now();
475
476        // From RFC 6762 section 5.2:
477        // "... The querier should plan to issue a query at 80% of the record
478        // lifetime, and then if no answer is received, at 85%, 90%, and 95%."
479        let refresh = get_expiration_time(created, ttl, 80);
480
481        let expires = get_expiration_time(created, ttl, 100);
482
483        Self {
484            entry: DnsEntry::new(name.to_string(), ty, class),
485            ttl,
486            created,
487            expires,
488            refresh,
489            new_name: None,
490        }
491    }
492
493    pub const fn get_ttl(&self) -> u32 {
494        self.ttl
495    }
496
497    pub const fn get_expire_time(&self) -> Instant {
498        self.expires
499    }
500
501    pub const fn get_refresh_time(&self) -> Instant {
502        self.refresh
503    }
504
505    pub fn is_expired(&self, now: Instant) -> bool {
506        now >= self.expires
507    }
508
509    /// Returns whether record expires in 1 second.
510    ///
511    /// This is useful because mDNS sets TTL to 1 (not 0) for expiring records.
512    pub fn expires_soon(&self, now: Instant) -> bool {
513        now + Duration::from_millis(1000) >= self.expires
514    }
515
516    pub fn refresh_due(&self, now: Instant) -> bool {
517        now >= self.refresh
518    }
519
520    /// Returns whether `now` (in millis) has passed half of TTL.
521    pub fn halflife_passed(&self, now: Instant) -> bool {
522        let halflife = get_expiration_time(self.created, self.ttl, 50);
523        now > halflife
524    }
525
526    pub fn is_unique(&self) -> bool {
527        self.entry.cache_flush
528    }
529
530    /// Updates the refresh time to be the same as the expire time so that
531    /// this record will not refresh again and will just expire.
532    pub fn refresh_no_more(&mut self) {
533        self.refresh = get_expiration_time(self.created, self.ttl, 100);
534    }
535
536    /// Returns if this record is due for refresh. If yes, `refresh` time is updated.
537    pub fn refresh_maybe(&mut self, now: Instant) -> bool {
538        if self.is_expired(now) || !self.refresh_due(now) {
539            return false;
540        }
541
542        trace!(
543            "{} qtype {} is due to refresh",
544            &self.entry.name,
545            self.entry.ty
546        );
547
548        // From RFC 6762 section 5.2:
549        // "... The querier should plan to issue a query at 80% of the record
550        // lifetime, and then if no answer is received, at 85%, 90%, and 95%."
551        //
552        // If the answer is received in time, 'refresh' will be reset outside
553        // this function, back to 80% of the new TTL.
554        if self.refresh == get_expiration_time(self.created, self.ttl, 80) {
555            self.refresh = get_expiration_time(self.created, self.ttl, 85);
556        } else if self.refresh == get_expiration_time(self.created, self.ttl, 85) {
557            self.refresh = get_expiration_time(self.created, self.ttl, 90);
558        } else if self.refresh == get_expiration_time(self.created, self.ttl, 90) {
559            self.refresh = get_expiration_time(self.created, self.ttl, 95);
560        } else {
561            self.refresh_no_more();
562        }
563
564        true
565    }
566
567    /// Return the absolute time for this record being created
568    pub const fn get_created(&self) -> Instant {
569        self.created
570    }
571
572    /// Set the expiration time
573    fn set_expire(&mut self, expire_at: Instant) {
574        self.expires = expire_at;
575    }
576
577    /// Moves this record's timestamps back by `elapsed`, as if `elapsed`
578    /// more time had passed since it was created.
579    ///
580    /// If a timestamp cannot be moved back that far, the record is marked
581    /// expired instead.
582    fn age_by(&mut self, elapsed: Duration) {
583        match (
584            self.created.checked_sub(elapsed),
585            self.expires.checked_sub(elapsed),
586            self.refresh.checked_sub(elapsed),
587        ) {
588            (Some(created), Some(expires), Some(refresh)) => {
589                self.created = created;
590                self.expires = expires;
591                self.refresh = refresh;
592            }
593            _ => {
594                // `created` is in the past, so the record is expired now.
595                self.expires = self.created;
596                self.refresh = self.created;
597            }
598        }
599    }
600
601    fn reset_ttl(&mut self, other: &Self) {
602        self.ttl = other.ttl;
603        self.created = other.created;
604        self.expires = get_expiration_time(self.created, self.ttl, 100);
605        self.refresh = if self.ttl > 1 {
606            get_expiration_time(self.created, self.ttl, 80)
607        } else {
608            // If TTL is 1, it means this record is expiring,
609            // then we set refresh to the same time as expires.
610            self.expires
611        };
612    }
613
614    /// Modify TTL to reflect the remaining life time from `now`.
615    pub fn update_ttl(&mut self, now: Instant) {
616        let elapsed = now.saturating_duration_since(self.created);
617        self.ttl = self.ttl.saturating_sub(elapsed.as_secs() as u32);
618    }
619
620    pub fn set_new_name(&mut self, new_name: String) {
621        if new_name == self.entry.name {
622            self.new_name = None;
623        } else {
624            self.new_name = Some(new_name);
625        }
626    }
627
628    pub fn get_new_name(&self) -> Option<&str> {
629        self.new_name.as_deref()
630    }
631
632    /// Return the new name if exists, otherwise the regular name in DnsEntry.
633    pub(crate) fn get_name(&self) -> &str {
634        self.new_name.as_deref().unwrap_or(&self.entry.name)
635    }
636
637    pub fn get_original_name(&self) -> &str {
638        &self.entry.name
639    }
640}
641
642impl PartialEq for DnsRecord {
643    fn eq(&self, other: &Self) -> bool {
644        self.entry == other.entry
645    }
646}
647
648/// Common methods for DNS resource records.
649pub trait DnsRecordExt: fmt::Debug {
650    fn get_record(&self) -> &DnsRecord;
651    fn get_record_mut(&mut self) -> &mut DnsRecord;
652    /// Writes the rdata of this record into `packet`.
653    fn write(&self, packet: &mut DnsOutPacket) -> WriteResult;
654    fn any(&self) -> &dyn Any;
655
656    /// Returns whether `other` record is considered the same except TTL.
657    fn matches(&self, other: &dyn DnsRecordExt) -> bool;
658
659    /// Returns whether `other` record has the same rdata.
660    fn rrdata_match(&self, other: &dyn DnsRecordExt) -> bool;
661
662    /// Returns the result based on a byte-level comparison of `rdata`.
663    /// If `other` is not valid, returns `Greater`.
664    fn compare_rdata(&self, other: &dyn DnsRecordExt) -> cmp::Ordering;
665
666    /// Returns the result based on "lexicographically later" defined below.
667    fn compare(&self, other: &dyn DnsRecordExt) -> cmp::Ordering {
668        /*
669        RFC 6762: https://datatracker.ietf.org/doc/html/rfc6762#section-8.2
670
671        ... The determination of "lexicographically later" is performed by first
672        comparing the record class (excluding the cache-flush bit described
673        in Section 10.2), then the record type, then raw comparison of the
674        binary content of the rdata without regard for meaning or structure.
675        If the record classes differ, then the numerically greater class is
676        considered "lexicographically later".  Otherwise, if the record types
677        differ, then the numerically greater type is considered
678        "lexicographically later".  If the rrtype and rrclass both match,
679        then the rdata is compared. ...
680        */
681        match self.get_class().cmp(&other.get_class()) {
682            cmp::Ordering::Equal => match self.get_type().cmp(&other.get_type()) {
683                cmp::Ordering::Equal => self.compare_rdata(other),
684                not_equal => not_equal,
685            },
686            not_equal => not_equal,
687        }
688    }
689
690    /// Returns a human-readable string of rdata.
691    fn rdata_print(&self) -> String;
692
693    /// Returns the class only, excluding class_flush / unique bit.
694    fn get_class(&self) -> u16 {
695        self.get_record().entry.class
696    }
697
698    fn get_cache_flush(&self) -> bool {
699        self.get_record().entry.cache_flush
700    }
701
702    /// Return the new name if exists, otherwise the regular name in DnsEntry.
703    fn get_name(&self) -> &str {
704        self.get_record().get_name()
705    }
706
707    fn get_type(&self) -> RRType {
708        self.get_record().entry.ty
709    }
710
711    /// Resets TTL using `other` record.
712    /// `self.refresh` and `self.expires` are also reset.
713    fn reset_ttl(&mut self, other: &dyn DnsRecordExt) {
714        self.get_record_mut().reset_ttl(other.get_record());
715    }
716
717    fn get_created(&self) -> Instant {
718        self.get_record().get_created()
719    }
720
721    fn get_expire(&self) -> Instant {
722        self.get_record().get_expire_time()
723    }
724
725    fn set_expire(&mut self, expire_at: Instant) {
726        self.get_record_mut().set_expire(expire_at);
727    }
728
729    /// Set expire as `expire_at` if it is sooner than the current `expire`.
730    fn set_expire_sooner(&mut self, expire_at: Instant) {
731        if expire_at < self.get_expire() {
732            self.get_record_mut().set_expire(expire_at);
733        }
734    }
735
736    /// Ages the record by `elapsed`.
737    fn age_by(&mut self, elapsed: Duration) {
738        self.get_record_mut().age_by(elapsed);
739    }
740
741    /// Returns true if the record expires in 1 second from `now`.
742    fn expires_soon(&self, now: Instant) -> bool {
743        self.get_record().expires_soon(now)
744    }
745
746    /// Given `now`, if the record is due to refresh, this method updates the refresh time
747    /// and returns the new refresh time. Otherwise, returns None.
748    fn updated_refresh_time(&mut self, now: Instant) -> Option<Instant> {
749        if self.get_record_mut().refresh_maybe(now) {
750            Some(self.get_record().get_refresh_time())
751        } else {
752            None
753        }
754    }
755
756    /// Returns true if another record has matched content,
757    /// and if its TTL is at least half of this record's.
758    fn suppressed_by_answer(&self, other: &dyn DnsRecordExt) -> bool {
759        self.matches(other) && (other.get_record().ttl > self.get_record().ttl / 2)
760    }
761
762    /// Required by RFC 6762 Section 7.1: Known-Answer Suppression.
763    fn suppressed_by(&self, msg: &DnsIncoming) -> bool {
764        for answer in msg.answers.iter() {
765            if self.suppressed_by_answer(answer.as_ref()) {
766                return true;
767            }
768        }
769        false
770    }
771
772    fn clone_box(&self) -> DnsRecordBox;
773
774    fn boxed(self) -> DnsRecordBox;
775}
776
777/// Resource Record for IPv4 address or IPv6 address.
778#[derive(Debug, Clone)]
779pub(crate) struct DnsAddress {
780    pub(crate) record: DnsRecord,
781    address: IpAddr,
782    pub(crate) interface_id: InterfaceId,
783}
784
785impl DnsAddress {
786    pub fn new(
787        name: &str,
788        ty: RRType,
789        class: u16,
790        ttl: u32,
791        address: IpAddr,
792        interface_id: InterfaceId,
793    ) -> Self {
794        let record = DnsRecord::new(name, ty, class, ttl);
795        Self {
796            record,
797            address,
798            interface_id,
799        }
800    }
801
802    pub fn address(&self) -> ScopedIp {
803        match self.address {
804            IpAddr::V4(v4) => ScopedIp::V4(ScopedIpV4 {
805                addr: v4,
806                interface_ids: vec![self.interface_id.clone()],
807            }),
808            IpAddr::V6(v6) => ScopedIp::V6(ScopedIpV6 {
809                addr: v6,
810                scope_id: self.interface_id.clone(),
811            }),
812        }
813    }
814}
815
816impl DnsRecordExt for DnsAddress {
817    fn get_record(&self) -> &DnsRecord {
818        &self.record
819    }
820
821    fn get_record_mut(&mut self) -> &mut DnsRecord {
822        &mut self.record
823    }
824
825    fn write(&self, packet: &mut DnsOutPacket) -> WriteResult {
826        match self.address {
827            IpAddr::V4(addr) => packet.write_bytes(addr.octets().as_ref()),
828            IpAddr::V6(addr) => packet.write_bytes(addr.octets().as_ref()),
829        };
830        Ok(())
831    }
832
833    fn any(&self) -> &dyn Any {
834        self
835    }
836
837    fn matches(&self, other: &dyn DnsRecordExt) -> bool {
838        if let Some(other_a) = other.any().downcast_ref::<Self>() {
839            return self.address == other_a.address
840                && self.record.entry == other_a.record.entry
841                && self.interface_id == other_a.interface_id;
842        }
843        false
844    }
845
846    fn rrdata_match(&self, other: &dyn DnsRecordExt) -> bool {
847        if let Some(other_a) = other.any().downcast_ref::<Self>() {
848            return self.address == other_a.address;
849        }
850        false
851    }
852
853    fn compare_rdata(&self, other: &dyn DnsRecordExt) -> cmp::Ordering {
854        if let Some(other_a) = other.any().downcast_ref::<Self>() {
855            self.address.cmp(&other_a.address)
856        } else {
857            cmp::Ordering::Greater
858        }
859    }
860
861    fn rdata_print(&self) -> String {
862        format!("{}", self.address)
863    }
864
865    fn clone_box(&self) -> DnsRecordBox {
866        Box::new(self.clone())
867    }
868
869    fn boxed(self) -> DnsRecordBox {
870        Box::new(self)
871    }
872}
873
874/// Resource Record for a DNS pointer
875#[derive(Debug, Clone)]
876pub struct DnsPointer {
877    record: DnsRecord,
878    alias: String, // the full name of Service Instance
879}
880
881impl DnsPointer {
882    pub fn new(name: &str, ty: RRType, class: u16, ttl: u32, alias: String) -> Self {
883        let record = DnsRecord::new(name, ty, class, ttl);
884        Self { record, alias }
885    }
886
887    pub fn alias(&self) -> &str {
888        &self.alias
889    }
890}
891
892impl DnsRecordExt for DnsPointer {
893    fn get_record(&self) -> &DnsRecord {
894        &self.record
895    }
896
897    fn get_record_mut(&mut self) -> &mut DnsRecord {
898        &mut self.record
899    }
900
901    fn write(&self, packet: &mut DnsOutPacket) -> WriteResult {
902        packet.write_name(&self.alias)
903    }
904
905    fn any(&self) -> &dyn Any {
906        self
907    }
908
909    fn matches(&self, other: &dyn DnsRecordExt) -> bool {
910        if let Some(other_ptr) = other.any().downcast_ref::<Self>() {
911            return self.alias == other_ptr.alias && self.record.entry == other_ptr.record.entry;
912        }
913        false
914    }
915
916    fn rrdata_match(&self, other: &dyn DnsRecordExt) -> bool {
917        if let Some(other_ptr) = other.any().downcast_ref::<Self>() {
918            return self.alias == other_ptr.alias;
919        }
920        false
921    }
922
923    fn compare_rdata(&self, other: &dyn DnsRecordExt) -> cmp::Ordering {
924        if let Some(other_ptr) = other.any().downcast_ref::<Self>() {
925            self.alias.cmp(&other_ptr.alias)
926        } else {
927            cmp::Ordering::Greater
928        }
929    }
930
931    fn rdata_print(&self) -> String {
932        self.alias.clone()
933    }
934
935    fn clone_box(&self) -> DnsRecordBox {
936        Box::new(self.clone())
937    }
938
939    fn boxed(self) -> DnsRecordBox {
940        Box::new(self)
941    }
942}
943
944/// Resource Record for a DNS service.
945#[derive(Debug, Clone)]
946pub struct DnsSrv {
947    pub(crate) record: DnsRecord,
948    pub(crate) priority: u16, // lower number means higher priority. Should be 0 in common cases.
949    pub(crate) weight: u16,   // Should be 0 in common cases
950    host: String,
951    port: u16,
952}
953
954impl DnsSrv {
955    pub fn new(
956        name: &str,
957        class: u16,
958        ttl: u32,
959        priority: u16,
960        weight: u16,
961        port: u16,
962        host: String,
963    ) -> Self {
964        let record = DnsRecord::new(name, RRType::SRV, class, ttl);
965        Self {
966            record,
967            priority,
968            weight,
969            host,
970            port,
971        }
972    }
973
974    pub fn host(&self) -> &str {
975        &self.host
976    }
977
978    pub fn port(&self) -> u16 {
979        self.port
980    }
981
982    pub fn set_host(&mut self, host: String) {
983        self.host = host;
984    }
985}
986
987impl DnsRecordExt for DnsSrv {
988    fn get_record(&self) -> &DnsRecord {
989        &self.record
990    }
991
992    fn get_record_mut(&mut self) -> &mut DnsRecord {
993        &mut self.record
994    }
995
996    fn write(&self, packet: &mut DnsOutPacket) -> WriteResult {
997        packet.write_short(self.priority);
998        packet.write_short(self.weight);
999        packet.write_short(self.port);
1000        packet.write_name(&self.host)
1001    }
1002
1003    fn any(&self) -> &dyn Any {
1004        self
1005    }
1006
1007    fn matches(&self, other: &dyn DnsRecordExt) -> bool {
1008        if let Some(other_svc) = other.any().downcast_ref::<Self>() {
1009            return self.host == other_svc.host
1010                && self.port == other_svc.port
1011                && self.weight == other_svc.weight
1012                && self.priority == other_svc.priority
1013                && self.record.entry == other_svc.record.entry;
1014        }
1015        false
1016    }
1017
1018    fn rrdata_match(&self, other: &dyn DnsRecordExt) -> bool {
1019        if let Some(other_srv) = other.any().downcast_ref::<Self>() {
1020            return self.host == other_srv.host
1021                && self.port == other_srv.port
1022                && self.weight == other_srv.weight
1023                && self.priority == other_srv.priority;
1024        }
1025        false
1026    }
1027
1028    fn compare_rdata(&self, other: &dyn DnsRecordExt) -> cmp::Ordering {
1029        let Some(other_srv) = other.any().downcast_ref::<Self>() else {
1030            return cmp::Ordering::Greater;
1031        };
1032
1033        // 1. compare `priority`
1034        match self
1035            .priority
1036            .to_be_bytes()
1037            .cmp(&other_srv.priority.to_be_bytes())
1038        {
1039            cmp::Ordering::Equal => {
1040                // 2. compare `weight`
1041                match self
1042                    .weight
1043                    .to_be_bytes()
1044                    .cmp(&other_srv.weight.to_be_bytes())
1045                {
1046                    cmp::Ordering::Equal => {
1047                        // 3. compare `port`.
1048                        match self.port.to_be_bytes().cmp(&other_srv.port.to_be_bytes()) {
1049                            cmp::Ordering::Equal => self.host.cmp(&other_srv.host),
1050                            not_equal => not_equal,
1051                        }
1052                    }
1053                    not_equal => not_equal,
1054                }
1055            }
1056            not_equal => not_equal,
1057        }
1058    }
1059
1060    fn rdata_print(&self) -> String {
1061        format!(
1062            "priority: {}, weight: {}, port: {}, host: {}",
1063            self.priority, self.weight, self.port, self.host
1064        )
1065    }
1066
1067    fn clone_box(&self) -> DnsRecordBox {
1068        Box::new(self.clone())
1069    }
1070
1071    fn boxed(self) -> DnsRecordBox {
1072        Box::new(self)
1073    }
1074}
1075
1076/// Resource Record for a DNS TXT record.
1077///
1078/// From [RFC 6763 section 6]:
1079///
1080/// The format of each constituent string within the DNS TXT record is a
1081/// single length byte, followed by 0-255 bytes of text data.
1082///
1083/// DNS-SD uses DNS TXT records to store arbitrary key/value pairs
1084///    conveying additional information about the named service.  Each
1085///    key/value pair is encoded as its own constituent string within the
1086///    DNS TXT record, in the form "key=value" (without the quotation
1087///    marks).  Everything up to the first '=' character is the key (Section
1088///    6.4).  Everything after the first '=' character to the end of the
1089///    string (including subsequent '=' characters, if any) is the value
1090#[derive(Clone)]
1091pub struct DnsTxt {
1092    pub(crate) record: DnsRecord,
1093    text: Vec<u8>,
1094}
1095
1096impl DnsTxt {
1097    pub fn new(name: &str, class: u16, ttl: u32, text: Vec<u8>) -> Self {
1098        let record = DnsRecord::new(name, RRType::TXT, class, ttl);
1099        Self { record, text }
1100    }
1101
1102    pub fn text(&self) -> &[u8] {
1103        &self.text
1104    }
1105}
1106
1107impl DnsRecordExt for DnsTxt {
1108    fn get_record(&self) -> &DnsRecord {
1109        &self.record
1110    }
1111
1112    fn get_record_mut(&mut self) -> &mut DnsRecord {
1113        &mut self.record
1114    }
1115
1116    fn write(&self, packet: &mut DnsOutPacket) -> WriteResult {
1117        packet.write_bytes(&self.text);
1118        Ok(())
1119    }
1120
1121    fn any(&self) -> &dyn Any {
1122        self
1123    }
1124
1125    fn matches(&self, other: &dyn DnsRecordExt) -> bool {
1126        if let Some(other_txt) = other.any().downcast_ref::<Self>() {
1127            return self.text == other_txt.text && self.record.entry == other_txt.record.entry;
1128        }
1129        false
1130    }
1131
1132    fn rrdata_match(&self, other: &dyn DnsRecordExt) -> bool {
1133        if let Some(other_txt) = other.any().downcast_ref::<Self>() {
1134            return self.text == other_txt.text;
1135        }
1136        false
1137    }
1138
1139    fn compare_rdata(&self, other: &dyn DnsRecordExt) -> cmp::Ordering {
1140        if let Some(other_txt) = other.any().downcast_ref::<Self>() {
1141            self.text.cmp(&other_txt.text)
1142        } else {
1143            cmp::Ordering::Greater
1144        }
1145    }
1146
1147    fn rdata_print(&self) -> String {
1148        format!("{:?}", decode_txt(&self.text))
1149    }
1150
1151    fn clone_box(&self) -> DnsRecordBox {
1152        Box::new(self.clone())
1153    }
1154
1155    fn boxed(self) -> DnsRecordBox {
1156        Box::new(self)
1157    }
1158}
1159
1160impl fmt::Debug for DnsTxt {
1161    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
1162        let properties = decode_txt(&self.text);
1163        write!(
1164            f,
1165            "DnsTxt {{ record: {:?}, text: {:?} }}",
1166            self.record, properties
1167        )
1168    }
1169}
1170
1171/// A DNS host information record
1172#[derive(Debug, Clone)]
1173struct DnsHostInfo {
1174    record: DnsRecord,
1175    cpu: String,
1176    os: String,
1177}
1178
1179impl DnsHostInfo {
1180    fn new(name: &str, ty: RRType, class: u16, ttl: u32, cpu: String, os: String) -> Self {
1181        let record = DnsRecord::new(name, ty, class, ttl);
1182        Self { record, cpu, os }
1183    }
1184}
1185
1186impl DnsRecordExt for DnsHostInfo {
1187    fn get_record(&self) -> &DnsRecord {
1188        &self.record
1189    }
1190
1191    fn get_record_mut(&mut self) -> &mut DnsRecord {
1192        &mut self.record
1193    }
1194
1195    fn write(&self, packet: &mut DnsOutPacket) -> WriteResult {
1196        debug!("Writing HInfo: cpu {} os {}", &self.cpu, &self.os);
1197        packet.write_bytes(self.cpu.as_bytes());
1198        packet.write_bytes(self.os.as_bytes());
1199        Ok(())
1200    }
1201
1202    fn any(&self) -> &dyn Any {
1203        self
1204    }
1205
1206    fn matches(&self, other: &dyn DnsRecordExt) -> bool {
1207        if let Some(other_hinfo) = other.any().downcast_ref::<Self>() {
1208            return self.cpu == other_hinfo.cpu
1209                && self.os == other_hinfo.os
1210                && self.record.entry == other_hinfo.record.entry;
1211        }
1212        false
1213    }
1214
1215    fn rrdata_match(&self, other: &dyn DnsRecordExt) -> bool {
1216        if let Some(other_hinfo) = other.any().downcast_ref::<Self>() {
1217            return self.cpu == other_hinfo.cpu && self.os == other_hinfo.os;
1218        }
1219        false
1220    }
1221
1222    fn compare_rdata(&self, other: &dyn DnsRecordExt) -> cmp::Ordering {
1223        if let Some(other_hinfo) = other.any().downcast_ref::<Self>() {
1224            match self.cpu.cmp(&other_hinfo.cpu) {
1225                cmp::Ordering::Equal => self.os.cmp(&other_hinfo.os),
1226                ordering => ordering,
1227            }
1228        } else {
1229            cmp::Ordering::Greater
1230        }
1231    }
1232
1233    fn rdata_print(&self) -> String {
1234        format!("cpu: {}, os: {}", self.cpu, self.os)
1235    }
1236
1237    fn clone_box(&self) -> DnsRecordBox {
1238        Box::new(self.clone())
1239    }
1240
1241    fn boxed(self) -> DnsRecordBox {
1242        Box::new(self)
1243    }
1244}
1245
1246/// Resource Record for negative responses
1247///
1248/// [RFC4034 section 4.1](https://datatracker.ietf.org/doc/html/rfc4034#section-4.1)
1249/// and
1250/// [RFC6762 section 6.1](https://datatracker.ietf.org/doc/html/rfc6762#section-6.1)
1251#[derive(Debug, Clone)]
1252pub struct DnsNSec {
1253    record: DnsRecord,
1254    next_domain: String,
1255    type_bitmap: Vec<u8>,
1256}
1257
1258impl DnsNSec {
1259    pub fn new(
1260        name: &str,
1261        class: u16,
1262        ttl: u32,
1263        next_domain: String,
1264        type_bitmap: Vec<u8>,
1265    ) -> Self {
1266        let record = DnsRecord::new(name, RRType::NSEC, class, ttl);
1267        Self {
1268            record,
1269            next_domain,
1270            type_bitmap,
1271        }
1272    }
1273
1274    /// Returns the types marked by `type_bitmap`
1275    pub fn _types(&self) -> Vec<u16> {
1276        // From RFC 4034: 4.1.2 The Type Bit Maps Field
1277        // https://datatracker.ietf.org/doc/html/rfc4034#section-4.1.2
1278        //
1279        // Each bitmap encodes the low-order 8 bits of RR types within the
1280        // window block, in network bit order.  The first bit is bit 0.  For
1281        // window block 0, bit 1 corresponds to RR type 1 (A), bit 2 corresponds
1282        // to RR type 2 (NS), and so forth.
1283
1284        let mut bit_num = 0;
1285        let mut results = Vec::new();
1286
1287        for byte in self.type_bitmap.iter() {
1288            let mut bit_mask: u8 = 0x80; // for bit 0 in network bit order
1289
1290            // check every bit in this byte, one by one.
1291            for _ in 0..8 {
1292                if (byte & bit_mask) != 0 {
1293                    results.push(bit_num);
1294                }
1295                bit_num += 1;
1296                bit_mask >>= 1; // mask for the next bit
1297            }
1298        }
1299        results
1300    }
1301}
1302
1303impl DnsRecordExt for DnsNSec {
1304    fn get_record(&self) -> &DnsRecord {
1305        &self.record
1306    }
1307
1308    fn get_record_mut(&mut self) -> &mut DnsRecord {
1309        &mut self.record
1310    }
1311
1312    fn write(&self, packet: &mut DnsOutPacket) -> WriteResult {
1313        packet.write_name(&self.next_domain)?;
1314        // The stored bitmap excludes the RFC 4034 window number and length.
1315        packet.write_bytes(&[0, self.type_bitmap.len() as u8]);
1316        packet.write_bytes(&self.type_bitmap);
1317        Ok(())
1318    }
1319
1320    fn any(&self) -> &dyn Any {
1321        self
1322    }
1323
1324    fn matches(&self, other: &dyn DnsRecordExt) -> bool {
1325        if let Some(other_record) = other.any().downcast_ref::<Self>() {
1326            return self.next_domain == other_record.next_domain
1327                && self.type_bitmap == other_record.type_bitmap
1328                && self.record.entry == other_record.record.entry;
1329        }
1330        false
1331    }
1332
1333    fn rrdata_match(&self, other: &dyn DnsRecordExt) -> bool {
1334        if let Some(other_record) = other.any().downcast_ref::<Self>() {
1335            return self.next_domain == other_record.next_domain
1336                && self.type_bitmap == other_record.type_bitmap;
1337        }
1338        false
1339    }
1340
1341    fn compare_rdata(&self, other: &dyn DnsRecordExt) -> cmp::Ordering {
1342        if let Some(other_nsec) = other.any().downcast_ref::<Self>() {
1343            match self.next_domain.cmp(&other_nsec.next_domain) {
1344                cmp::Ordering::Equal => self.type_bitmap.cmp(&other_nsec.type_bitmap),
1345                ordering => ordering,
1346            }
1347        } else {
1348            cmp::Ordering::Greater
1349        }
1350    }
1351
1352    fn rdata_print(&self) -> String {
1353        format!(
1354            "next_domain: {}, type_bitmap len: {}",
1355            self.next_domain,
1356            self.type_bitmap.len()
1357        )
1358    }
1359
1360    fn clone_box(&self) -> DnsRecordBox {
1361        Box::new(self.clone())
1362    }
1363
1364    fn boxed(self) -> DnsRecordBox {
1365        Box::new(self)
1366    }
1367}
1368
1369/// Which section of a DNS message an item belongs to.
1370#[derive(Clone, Copy, Debug)]
1371enum Section {
1372    Question,
1373    Answer,
1374    Authority,
1375    Additional,
1376}
1377
1378/// A single packet for outgoing DNS message.
1379pub struct DnsOutPacket {
1380    /// All bytes in `data` is the actual packet on the wire.
1381    data: Vec<u8>,
1382
1383    /// k: name, v: offset
1384    names: HashMap<String, u16>,
1385
1386    /// Max byte size of `data`. i.e. the max packet size.
1387    max_size: usize,
1388
1389    /// How many items `data` holds in each section, i.e. the header counts.
1390    question_count: u16,
1391    answer_count: u16,
1392    auth_count: u16,
1393    addi_count: u16,
1394}
1395
1396impl DnsOutPacket {
1397    fn new(max_size: usize) -> Self {
1398        Self {
1399            data: vec![0; MSG_HEADER_LEN],
1400            names: HashMap::new(),
1401            max_size,
1402            question_count: 0,
1403            answer_count: 0,
1404            auth_count: 0,
1405            addi_count: 0,
1406        }
1407    }
1408
1409    pub fn size(&self) -> usize {
1410        self.data.len()
1411    }
1412
1413    pub fn as_bytes(&self) -> &[u8] {
1414        &self.data
1415    }
1416
1417    /// True if nothing has been written into this packet yet.
1418    fn is_empty(&self) -> bool {
1419        self.question_count == 0
1420            && self.answer_count == 0
1421            && self.auth_count == 0
1422            && self.addi_count == 0
1423    }
1424
1425    /// Counts one more item in `section`.
1426    fn bump(&mut self, section: Section) {
1427        match section {
1428            Section::Question => self.question_count += 1,
1429            Section::Answer => self.answer_count += 1,
1430            Section::Authority => self.auth_count += 1,
1431            Section::Additional => self.addi_count += 1,
1432        }
1433    }
1434
1435    fn write_question(&mut self, question: &DnsQuestion) -> WriteResult {
1436        let start_size = self.size();
1437
1438        self.write_name(&question.entry.name).map_err(|e| {
1439            self.rollback(start_size);
1440            e
1441        })?;
1442        self.write_short(question.entry.ty as u16);
1443        self.write_short(question.entry.class);
1444
1445        if self.size() > self.max_size {
1446            self.rollback(start_size);
1447            return Err(WriteError::PacketFull);
1448        }
1449
1450        Ok(())
1451    }
1452
1453    /// Discards everything written since `start_size`, including the name
1454    /// compression offsets that point into the discarded bytes.
1455    fn rollback(&mut self, start_size: usize) {
1456        self.data.truncate(start_size);
1457        self.names
1458            .retain(|_, offset| (*offset as usize) < start_size);
1459    }
1460
1461    /// Writes a record (answer, authoritative answer, additional).
1462    ///
1463    /// In error cases nothing is written to the packet.
1464    fn write_record(&mut self, record_ext: &dyn DnsRecordExt) -> WriteResult {
1465        let start_size = self.size();
1466
1467        let record = record_ext.get_record();
1468        self.write_name(record.get_name())?;
1469        self.write_short(record.entry.ty as u16);
1470        if record.entry.cache_flush {
1471            // check "multicast"
1472            self.write_short(record.entry.class | CLASS_CACHE_FLUSH);
1473        } else {
1474            self.write_short(record.entry.class);
1475        }
1476
1477        self.write_u32(record.ttl);
1478
1479        // Placeholder for record size
1480        self.write_short(0);
1481        let record_offset = self.size();
1482
1483        if let Err(e) = record_ext.write(self) {
1484            self.rollback(start_size);
1485            return Err(e);
1486        }
1487
1488        self.set_short_at(record_offset - 2, (self.size() - record_offset) as u16);
1489
1490        if self.size() > self.max_size {
1491            self.rollback(start_size);
1492            return Err(WriteError::PacketFull);
1493        }
1494
1495        Ok(())
1496    }
1497
1498    fn set_short_at(&mut self, index: usize, value: u16) {
1499        self.data[index..index + 2].copy_from_slice(&value.to_be_bytes());
1500    }
1501
1502    /// Parses a DNS name that may contain escaped characters according to RFC 6763 Section 4.3.
1503    /// Returns a vector of labels where each label is the unescaped content.
1504    ///
1505    /// Escape sequences:
1506    /// - \\. becomes . (literal dot)
1507    /// - \\\\ becomes \\ (literal backslash)
1508    fn parse_escaped_name(name: &str) -> Vec<String> {
1509        let mut labels = Vec::new();
1510        let mut current_label = String::new();
1511        let mut chars = name.chars().peekable();
1512
1513        while let Some(ch) = chars.next() {
1514            match ch {
1515                '\\' => {
1516                    // Backslash escape sequence
1517                    if let Some(&next_ch) = chars.peek() {
1518                        match next_ch {
1519                            '.' | '\\' => {
1520                                // \\. or \\\\ - consume the backslash and add the escaped char
1521                                chars.next();
1522                                current_label.push(next_ch);
1523                            }
1524                            _ => {
1525                                // Not a recognized escape - treat backslash literally
1526                                current_label.push(ch);
1527                            }
1528                        }
1529                    } else {
1530                        // Trailing backslash - add it literally
1531                        current_label.push(ch);
1532                    }
1533                }
1534                '.' => {
1535                    // Unescaped dot - label separator
1536                    if !current_label.is_empty() {
1537                        labels.push(current_label.clone());
1538                        current_label.clear();
1539                    }
1540                }
1541                _ => {
1542                    current_label.push(ch);
1543                }
1544            }
1545        }
1546
1547        // Add the last label if not empty
1548        if !current_label.is_empty() {
1549            labels.push(current_label);
1550        }
1551
1552        labels
1553    }
1554
1555    // Write name to packet
1556    //
1557    // [RFC1035]
1558    // 4.1.4. Message compression
1559    //
1560    // In order to reduce the size of messages, the domain system utilizes a
1561    // compression scheme which eliminates the repetition of domain names in a
1562    // message.  In this scheme, an entire domain name or a list of labels at
1563    // the end of a domain name is replaced with a pointer to a prior occurrence
1564    // of the same name.
1565    // The pointer takes the form of a two octet sequence:
1566    //     +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+
1567    //     | 1  1|                OFFSET                   |
1568    //     +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+
1569    // The first two bits are ones.  This allows a pointer to be distinguished
1570    // from a label, since the label must begin with two zero bits because
1571    // labels are restricted to 63 octets or less.  (The 10 and 01 combinations
1572    // are reserved for future use.)  The OFFSET field specifies an offset from
1573    // the start of the message (i.e., the first octet of the ID field in the
1574    // domain header).  A zero offset specifies the first byte of the ID field,
1575    // etc.
1576    //
1577    // This function also handles RFC 6763 Section 4.3 escaping where dots and backslashes
1578    // in instance names are escaped (e.g., "My\\.Service" represents a single label "My.Service").
1579    // The actual name sent over the wire is the unescaped version.
1580    fn write_name(&mut self, name: &str) -> WriteResult {
1581        // Remove trailing dot if present
1582        let name_to_parse = name.strip_suffix('.').unwrap_or(name);
1583
1584        // Parse the name considering escape sequences
1585        let labels = Self::parse_escaped_name(name_to_parse);
1586
1587        if labels.is_empty() {
1588            self.write_byte(0);
1589            return Ok(());
1590        }
1591
1592        // Validate before writing anything.
1593        if labels.iter().any(|label| label.len() > MAX_LABEL_BYTES) {
1594            return Err(WriteError::NameTooLong);
1595        }
1596
1597        // Write each label
1598        for (i, label) in labels.iter().enumerate() {
1599            // Build the remaining name for compression (with dots as separators)
1600            let remaining: String = labels[i..].join(".");
1601
1602            // Check if we can use compression for the remaining part
1603            const POINTER_MASK: u16 = 0xC000;
1604            if let Some(&offset) = self.names.get(&remaining) {
1605                let pointer = offset | POINTER_MASK;
1606                self.write_short(pointer);
1607                return Ok(());
1608            }
1609
1610            // Store this position for potential future compression
1611            self.names.insert(remaining, self.size() as u16);
1612
1613            // Write the label
1614            self.write_utf8(label)?;
1615        }
1616
1617        // Write terminating zero byte
1618        self.write_byte(0);
1619        Ok(())
1620    }
1621
1622    fn write_byte(&mut self, v: u8) {
1623        self.data.push(v);
1624    }
1625
1626    fn write_bytes(&mut self, s: &[u8]) {
1627        self.data.extend(s);
1628    }
1629
1630    /// Writes a single label. Nothing is written if the label is too long to
1631    /// be encoded.
1632    fn write_utf8(&mut self, s: &str) -> WriteResult {
1633        if s.len() > MAX_LABEL_BYTES {
1634            return Err(WriteError::NameTooLong);
1635        }
1636        self.write_byte(s.len() as u8);
1637        self.write_bytes(s.as_bytes());
1638        Ok(())
1639    }
1640
1641    fn write_u32(&mut self, v: u32) {
1642        self.data.extend(&v.to_be_bytes());
1643    }
1644
1645    fn write_short(&mut self, v: u16) {
1646        self.data.extend(&v.to_be_bytes());
1647    }
1648
1649    /// Marks this finished packet as truncated, i.e. the message continues in
1650    /// the next packet.
1651    fn set_truncated(&mut self) {
1652        let flags = u16::from_be_bytes([self.data[2], self.data[3]]);
1653        self.set_short_at(2, flags | FLAGS_TC);
1654    }
1655
1656    /// Writes the header fields and finish the packet.
1657    /// This function should be only called when finishing a packet.
1658    ///
1659    /// The header format is based on RFC 1035 section 4.1.1:
1660    /// https://datatracker.ietf.org/doc/html/rfc1035#section-4.1.1
1661    //
1662    //                                  1  1  1  1  1  1
1663    //    0  1  2  3  4  5  6  7  8  9  0  1  2  3  4  5
1664    //    +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+
1665    //    |                      ID                       |
1666    //    +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+
1667    //    |QR|   Opcode  |AA|TC|RD|RA|   Z    |   RCODE   |
1668    //    +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+
1669    //    |                    QDCOUNT                    |
1670    //    +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+
1671    //    |                    ANCOUNT                    |
1672    //    +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+
1673    //    |                    NSCOUNT                    |
1674    //    +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+
1675    //    |                    ARCOUNT                    |
1676    //    +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+
1677    //
1678    fn write_header(&mut self, id: u16, flags: u16) {
1679        self.set_short_at(0, id);
1680        self.set_short_at(2, flags);
1681        self.set_short_at(4, self.question_count);
1682        self.set_short_at(6, self.answer_count);
1683        self.set_short_at(8, self.auth_count);
1684        self.set_short_at(10, self.addi_count);
1685    }
1686}
1687
1688/// Encodes a [`DnsOutgoing`] into one or more [`DnsOutPacket`], starting a new
1689/// packet whenever the current one runs out of room.
1690struct PacketBuilder<'a> {
1691    out: &'a DnsOutgoing,
1692
1693    /// Max size of a packet that holds more than one record.
1694    max_size: usize,
1695
1696    /// IP version these packets are bound for, which decides their absolute
1697    /// ceiling: see [`max_pkt_absolute`].
1698    is_ipv4: bool,
1699
1700    finished: Vec<DnsOutPacket>,
1701    current: DnsOutPacket,
1702}
1703
1704impl<'a> PacketBuilder<'a> {
1705    fn new(out: &'a DnsOutgoing, max_size: usize, is_ipv4: bool) -> Self {
1706        Self {
1707            out,
1708            max_size,
1709            is_ipv4,
1710            finished: Vec::new(),
1711            current: DnsOutPacket::new(max_size),
1712        }
1713    }
1714
1715    /// Writes one question or record into the current packet, starting a new
1716    /// packet if it does not fit in the current one.
1717    ///
1718    /// An item that cannot be encoded at all is skipped, leaving the packet as
1719    /// it was. Sections are written in message order, so an item that spills
1720    /// never lands ahead of one already written.
1721    fn add<F>(&mut self, section: Section, write: F)
1722    where
1723        F: Fn(&mut DnsOutPacket) -> WriteResult,
1724    {
1725        match write(&mut self.current) {
1726            Ok(()) => {
1727                self.current.bump(section);
1728                return;
1729            }
1730            // The item can never be encoded: skip it.
1731            Err(WriteError::NameTooLong) => return,
1732            Err(WriteError::PacketFull) => {}
1733        }
1734
1735        // Packet is full. Flush the current and create a new one.
1736        if !self.current.is_empty() {
1737            self.flush();
1738
1739            match write(&mut self.current) {
1740                Ok(()) => {
1741                    self.current.bump(section);
1742                    return;
1743                }
1744                Err(WriteError::NameTooLong) => return,
1745                Err(WriteError::PacketFull) => {}
1746            }
1747        }
1748
1749        // Packet is still full. A question such big is not legitimate.
1750        if matches!(section, Section::Question) {
1751            return;
1752        }
1753
1754        // Packet is still full. We will send this single record.
1755
1756        // RFC 6762 section 17:
1757        // "a record too large for one MTU-sized packet SHOULD be sent alone, in a
1758        // single IP datagram".
1759        self.current.max_size = max_pkt_absolute(self.is_ipv4);
1760
1761        if write(&mut self.current).is_ok() {
1762            self.current.bump(section);
1763            self.flush();
1764        } else {
1765            // Too big even for the hard ceiling: skip the record and carry on.
1766            self.current.max_size = self.max_size;
1767            debug!(
1768                "Record too big for absolute max size, skipping: {:?}",
1769                section
1770            );
1771        }
1772    }
1773
1774    /// Finishes the current packet and starts a new empty one.
1775    fn flush(&mut self) {
1776        self.current
1777            .write_header(self.out.wire_id(), self.out.flags);
1778
1779        let next = DnsOutPacket::new(self.max_size);
1780        self.finished
1781            .push(std::mem::replace(&mut self.current, next));
1782    }
1783
1784    fn finish(mut self) -> Vec<DnsOutPacket> {
1785        // Always produce at least one packet, even an empty one, but never leave a
1786        // trailing empty packet behind a full one.
1787        if !self.current.is_empty() || self.finished.is_empty() {
1788            self.flush();
1789        }
1790
1791        let mut packets = self.finished;
1792
1793        /*
1794        RFC 6762 section 7.2: https://datatracker.ietf.org/doc/html/rfc6762#section-7.2
1795        ...
1796            When a Multicast DNS querier sends a query to which it already knows some
1797            answers, it ... sets the TC (Truncated) bit in the header ... [so that the
1798            responder knows] to wait for the remaining known answers before responding.
1799         */
1800        if self.out.is_query() {
1801            if let Some((_last, rest)) = packets.split_last_mut() {
1802                for packet in rest {
1803                    packet.set_truncated();
1804                }
1805            }
1806        }
1807
1808        packets
1809    }
1810}
1811
1812/// Representation of one outgoing DNS message that could be sent in one or more packet(s).
1813#[derive(Debug)]
1814pub struct DnsOutgoing {
1815    flags: u16,
1816    id: u16,
1817    multicast: bool,
1818    questions: Vec<DnsQuestion>,
1819    answers: Vec<DnsRecordBox>,
1820    authorities: Vec<DnsRecordBox>,
1821    additionals: Vec<DnsRecordBox>,
1822    known_answer_count: i64, // for internal maintenance only
1823}
1824
1825impl DnsOutgoing {
1826    pub fn new(flags: u16) -> Self {
1827        Self {
1828            flags,
1829            id: 0,
1830            multicast: true,
1831            questions: Vec::new(),
1832            answers: Vec::new(),
1833            authorities: Vec::new(),
1834            additionals: Vec::new(),
1835            known_answer_count: 0,
1836        }
1837    }
1838
1839    pub fn questions(&self) -> &[DnsQuestion] {
1840        &self.questions
1841    }
1842
1843    /// For testing purposes only.
1844    pub(crate) fn _answers(&self) -> &[DnsRecordBox] {
1845        &self.answers
1846    }
1847
1848    pub fn answers_count(&self) -> usize {
1849        self.answers.len()
1850    }
1851
1852    pub fn authorities(&self) -> &[DnsRecordBox] {
1853        &self.authorities
1854    }
1855
1856    pub fn additionals(&self) -> &[DnsRecordBox] {
1857        &self.additionals
1858    }
1859
1860    pub fn known_answer_count(&self) -> i64 {
1861        self.known_answer_count
1862    }
1863
1864    pub fn set_id(&mut self, id: u16) {
1865        self.id = id;
1866    }
1867
1868    /// Marks whether this message is destined for multicast (the default) or unicast.
1869    pub fn set_multicast(&mut self, multicast: bool) {
1870        self.multicast = multicast;
1871    }
1872
1873    /// The id to put in the header, always 0 for multicast.
1874    const fn wire_id(&self) -> u16 {
1875        if self.multicast {
1876            0
1877        } else {
1878            // RFC 6762 ยง6.7: a legacy unicast response MUST echo the
1879            // querier's message id.
1880            self.id
1881        }
1882    }
1883
1884    pub const fn is_query(&self) -> bool {
1885        (self.flags & FLAGS_QR_MASK) == FLAGS_QR_QUERY
1886    }
1887
1888    // Adds an additional answer
1889
1890    // From: RFC 6763, DNS-Based Service Discovery, February 2013
1891
1892    // 12.  DNS Additional Record Generation
1893
1894    //    DNS has an efficiency feature whereby a DNS server may place
1895    //    additional records in the additional section of the DNS message.
1896    //    These additional records are records that the client did not
1897    //    explicitly request, but the server has reasonable grounds to expect
1898    //    that the client might request them shortly, so including them can
1899    //    save the client from having to issue additional queries.
1900
1901    //    This section recommends which additional records SHOULD be generated
1902    //    to improve network efficiency, for both Unicast and Multicast DNS-SD
1903    //    responses.
1904
1905    // 12.1.  PTR Records
1906
1907    //    When including a DNS-SD Service Instance Enumeration or Selective
1908    //    Instance Enumeration (subtype) PTR record in a response packet, the
1909    //    server/responder SHOULD include the following additional records:
1910
1911    //    o  The SRV record(s) named in the PTR rdata.
1912    //    o  The TXT record(s) named in the PTR rdata.
1913    //    o  All address records (type "A" and "AAAA") named in the SRV rdata.
1914
1915    // 12.2.  SRV Records
1916
1917    //    When including an SRV record in a response packet, the
1918    //    server/responder SHOULD include the following additional records:
1919
1920    //    o  All address records (type "A" and "AAAA") named in the SRV rdata.
1921    pub fn add_additional_answer(&mut self, answer: impl DnsRecordExt + 'static) {
1922        trace!("add_additional_answer: {:?}", &answer);
1923        self.additionals.push(answer.boxed());
1924    }
1925
1926    /// A workaround as Rust doesn't allow us to pass DnsRecordBox in as `impl DnsRecordExt`
1927    pub fn add_answer_box(&mut self, answer_box: DnsRecordBox) {
1928        self.answers.push(answer_box);
1929    }
1930
1931    pub fn add_authority(&mut self, record: DnsRecordBox) {
1932        self.authorities.push(record);
1933    }
1934
1935    /// Retains only the answers for which `keep` returns true.
1936    pub(crate) fn retain_answers<F>(&mut self, mut keep: F)
1937    where
1938        F: FnMut(&DnsRecordBox) -> bool,
1939    {
1940        self.answers.retain(|record| keep(record));
1941    }
1942
1943    /// Retains only the additional records for which `keep` returns true.
1944    pub(crate) fn retain_additionals<F>(&mut self, mut keep: F)
1945    where
1946        F: FnMut(&DnsRecordBox) -> bool,
1947    {
1948        self.additionals.retain(|record| keep(record));
1949    }
1950
1951    /// Returns true if `answer` is added to the outgoing msg.
1952    /// Returns false if `answer` was not added as it is suppressed by the incoming `msg`.
1953    pub fn add_answer(
1954        &mut self,
1955        msg: &DnsIncoming,
1956        answer: impl DnsRecordExt + Send + 'static,
1957    ) -> bool {
1958        trace!("Check for add_answer");
1959        if answer.suppressed_by(msg) {
1960            trace!("my answer is suppressed by incoming msg");
1961            self.known_answer_count += 1;
1962            return false;
1963        }
1964
1965        self.add_answer_record(answer);
1966        true
1967    }
1968
1969    /// Adds `answer` to the outgoing msg unconditionally. Its full TTL is
1970    /// written on the wire.
1971    pub fn add_answer_record(&mut self, answer: impl DnsRecordExt + Send + 'static) {
1972        trace!("add_answer push: {:?}", &answer);
1973        self.answers.push(answer.boxed());
1974    }
1975
1976    /// Adds a PTR answer for `service` along with recommended additional records
1977    /// (SRV, TXT, and address records) per [RFC 6763 Section 12.1].
1978    ///
1979    /// Resolves any name conflicts via `dns_registry` and selects addresses
1980    /// matching the given interface. Does nothing if no addresses are available
1981    /// on `intf` or if the PTR answer is suppressed by known-answer entries in `msg`.
1982    ///
1983    /// [RFC 6763 Section 12.1]: https://tools.ietf.org/html/rfc6763#section-12.1
1984    pub(crate) fn add_answer_with_additionals(
1985        &mut self,
1986        msg: &DnsIncoming,
1987        service: &ServiceInfo,
1988        intf: &MyIntf,
1989        dns_registry: &DnsRegistry,
1990        is_ipv4: bool,
1991    ) {
1992        let intf_addrs = if is_ipv4 {
1993            service.get_addrs_on_my_intf_v4(intf)
1994        } else {
1995            service.get_addrs_on_my_intf_v6(intf)
1996        };
1997        if intf_addrs.is_empty() {
1998            trace!("No addrs on LAN of intf {:?}", intf);
1999            return;
2000        }
2001
2002        // check if we changed our name due to conflicts.
2003        let service_fullname = dns_registry.resolve_name(service.get_fullname());
2004        let hostname = dns_registry.resolve_name(service.get_hostname());
2005
2006        let ptr_added = self.add_answer(
2007            msg,
2008            DnsPointer::new(
2009                service.get_type(),
2010                RRType::PTR,
2011                CLASS_IN,
2012                service.get_other_ttl(),
2013                service_fullname.to_string(),
2014            ),
2015        );
2016
2017        if !ptr_added {
2018            trace!("answer was not added for msg {:?}", msg);
2019            return;
2020        }
2021
2022        if let Some(sub) = service.get_subtype() {
2023            trace!("Adding subdomain {}", sub);
2024            self.add_additional_answer(DnsPointer::new(
2025                sub,
2026                RRType::PTR,
2027                CLASS_IN,
2028                service.get_other_ttl(),
2029                service_fullname.to_string(),
2030            ));
2031        }
2032
2033        // Add recommended additional answers according to
2034        // https://tools.ietf.org/html/rfc6763#section-12.1.
2035        self.add_additional_answer(DnsSrv::new(
2036            service_fullname,
2037            CLASS_IN | CLASS_CACHE_FLUSH,
2038            service.get_host_ttl(),
2039            service.get_priority(),
2040            service.get_weight(),
2041            service.get_port(),
2042            hostname.to_string(),
2043        ));
2044
2045        self.add_additional_answer(DnsTxt::new(
2046            service_fullname,
2047            CLASS_IN | CLASS_CACHE_FLUSH,
2048            service.get_other_ttl(),
2049            service.generate_txt(),
2050        ));
2051
2052        for address in intf_addrs {
2053            self.add_additional_answer(DnsAddress::new(
2054                hostname,
2055                ip_address_rr_type(&address),
2056                CLASS_IN | CLASS_CACHE_FLUSH,
2057                service.get_host_ttl(),
2058                address,
2059                intf.into(),
2060            ));
2061        }
2062    }
2063
2064    pub fn add_question(&mut self, name: &str, qtype: RRType) {
2065        let q = DnsQuestion {
2066            entry: DnsEntry::new(name.to_string(), qtype, CLASS_IN),
2067        };
2068        self.questions.push(q);
2069    }
2070
2071    /// Adjust records so the message is a valid legacy unicast response:
2072    ///
2073    /// - Clear the cache-flush (unique) bit: a legacy resolver
2074    ///   doesn't know about it and may misinterpret responses where it is set.
2075    /// - Cap the TTL at [`LEGACY_UNICAST_MAX_TTL`] seconds: legacy resolvers
2076    ///   cache records without the mDNS cache-coherency mechanisms, so the true
2077    ///   (longer) TTL must not leak out to them.
2078    ///
2079    /// Refer to [RFC 6762 Section 6.7] for details.
2080    pub fn update_records_for_legacy_unicast(&mut self) {
2081        let update = |rec: &mut DnsRecordBox| {
2082            let record = rec.get_record_mut();
2083            record.entry.cache_flush = false;
2084            record.ttl = record.ttl.min(LEGACY_UNICAST_MAX_TTL);
2085        };
2086        for rec in &mut self.answers {
2087            update(rec);
2088        }
2089        for rec in &mut self.additionals {
2090            update(rec);
2091        }
2092        for rec in &mut self.authorities {
2093            update(rec);
2094        }
2095    }
2096
2097    /// Returns a list of actual DNS packet data to be sent on the wire, each no
2098    /// bigger than `max_size`, over the IP version given by `is_ipv4`.
2099    ///
2100    /// Most callers want [`MAX_PKT_DEFAULT`] for `max_size`.
2101    pub fn to_data_on_wire(&self, max_size: usize, is_ipv4: bool) -> Vec<Vec<u8>> {
2102        let packet_list = self.to_packets(max_size, is_ipv4);
2103        packet_list.into_iter().map(|p| p.data).collect()
2104    }
2105
2106    /// Encode self into one or more packets, each no bigger than `max_size`.
2107    ///
2108    /// Questions and records are written in message order and spill into a new
2109    /// packet whenever the current one is full, so none is dropped for lack of
2110    /// room. The one exception is a single record too big to fit in an otherwise
2111    /// empty packet: it is sent alone in an oversized packet, per RFC 6762
2112    /// section 17.
2113    ///
2114    /// `is_ipv4` tells which IP version the packets are bound for, and so how big
2115    /// that lone oversized packet may get: see [`max_pkt_absolute`]. A record too
2116    /// big even for that could not be sent at all, and is dropped.
2117    ///
2118    /// `max_size` must be no bigger than [`MAX_PKT_ABSOLUTE_IPV6`], the RFC 6762
2119    /// section 17 ceiling that is legal over either IP version;
2120    /// [`ServiceDaemon::set_max_packet_size`](crate::ServiceDaemon::set_max_packet_size)
2121    /// caps what it accepts. Most callers want [`MAX_PKT_DEFAULT`].
2122    pub fn to_packets(&self, max_size: usize, is_ipv4: bool) -> Vec<DnsOutPacket> {
2123        debug_assert!(
2124            max_size <= MAX_PKT_ABSOLUTE_IPV6,
2125            "max_size {} exceeds the RFC 6762 section 17 ceiling",
2126            max_size
2127        );
2128        let mut builder = PacketBuilder::new(self, max_size, is_ipv4);
2129
2130        for question in self.questions.iter() {
2131            builder.add(Section::Question, |packet| packet.write_question(question));
2132        }
2133
2134        for answer in self.answers.iter() {
2135            builder.add(Section::Answer, |packet| {
2136                packet.write_record(answer.as_ref())
2137            });
2138        }
2139
2140        for auth in self.authorities.iter() {
2141            builder.add(Section::Authority, |packet| {
2142                packet.write_record(auth.as_ref())
2143            });
2144        }
2145
2146        for addi in self.additionals.iter() {
2147            builder.add(Section::Additional, |packet| {
2148                packet.write_record(addi.as_ref())
2149            });
2150        }
2151
2152        builder.finish()
2153    }
2154}
2155
2156/// An incoming DNS message. It could be a query or a response.
2157pub struct DnsIncoming {
2158    offset: usize,
2159    data: Vec<u8>,
2160    questions: Vec<DnsQuestion>,
2161    answers: Vec<DnsRecordBox>,
2162    authorities: Vec<DnsRecordBox>,
2163    additional: Vec<DnsRecordBox>,
2164    id: u16,
2165    flags: u16,
2166    num_questions: u16,
2167    num_answers: u16,
2168    num_authorities: u16,
2169    num_additionals: u16,
2170    interface_id: InterfaceId,
2171}
2172
2173/// Written by hand rather than derived, so we don't dump the raw packet unbounded.
2174impl fmt::Debug for DnsIncoming {
2175    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
2176        f.debug_struct("DnsIncoming")
2177            .field("offset", &self.offset)
2178            .field("questions", &self.questions)
2179            .field("answers", &self.answers)
2180            .field("authorities", &self.authorities)
2181            .field("additional", &self.additional)
2182            .field("id", &self.id)
2183            .field("flags", &self.flags)
2184            .field("num_questions", &self.num_questions)
2185            .field("num_answers", &self.num_answers)
2186            .field("num_authorities", &self.num_authorities)
2187            .field("num_additionals", &self.num_additionals)
2188            .field("interface_id", &self.interface_id)
2189            .finish()
2190    }
2191}
2192
2193impl DnsIncoming {
2194    pub fn new(data: Vec<u8>, interface_id: InterfaceId) -> Result<Self> {
2195        let mut incoming = Self {
2196            offset: 0,
2197            data,
2198            questions: Vec::new(),
2199            answers: Vec::new(),
2200            authorities: Vec::new(),
2201            additional: Vec::new(),
2202            id: 0,
2203            flags: 0,
2204            num_questions: 0,
2205            num_answers: 0,
2206            num_authorities: 0,
2207            num_additionals: 0,
2208            interface_id,
2209        };
2210
2211        /*
2212        RFC 1035 section 4.1: https://datatracker.ietf.org/doc/html/rfc1035#section-4.1
2213        ...
2214        All communications inside of the domain protocol are carried in a single
2215        format called a message.  The top level format of message is divided
2216        into 5 sections (some of which are empty in certain cases) shown below:
2217
2218            +---------------------+
2219            |        Header       |
2220            +---------------------+
2221            |       Question      | the question for the name server
2222            +---------------------+
2223            |        Answer       | RRs answering the question
2224            +---------------------+
2225            |      Authority      | RRs pointing toward an authority
2226            +---------------------+
2227            |      Additional     | RRs holding additional information
2228            +---------------------+
2229         */
2230        if let Err(e) = incoming.read_sections() {
2231            return Err(Error::Msg(format!(
2232                "{e}; raw packet length: {}",
2233                incoming.data.len(),
2234            )));
2235        }
2236
2237        Ok(incoming)
2238    }
2239
2240    /// Reads the five message sections in order. Kept separate from `new` so a
2241    /// parse failure can be annotated with the raw packet bytes.
2242    fn read_sections(&mut self) -> Result<()> {
2243        self.read_header()?;
2244        self.read_questions()?;
2245        self.read_answers()?;
2246        self.read_authorities()?;
2247        self.read_additional()?;
2248        Ok(())
2249    }
2250
2251    pub fn id(&self) -> u16 {
2252        self.id
2253    }
2254
2255    pub fn questions(&self) -> &[DnsQuestion] {
2256        &self.questions
2257    }
2258
2259    pub fn answers(&self) -> &[DnsRecordBox] {
2260        &self.answers
2261    }
2262
2263    pub fn authorities(&self) -> &[DnsRecordBox] {
2264        &self.authorities
2265    }
2266
2267    pub fn additionals(&self) -> &[DnsRecordBox] {
2268        &self.additional
2269    }
2270
2271    pub fn answers_mut(&mut self) -> &mut Vec<DnsRecordBox> {
2272        &mut self.answers
2273    }
2274
2275    pub fn authorities_mut(&mut self) -> &mut Vec<DnsRecordBox> {
2276        &mut self.authorities
2277    }
2278
2279    pub fn additionals_mut(&mut self) -> &mut Vec<DnsRecordBox> {
2280        &mut self.additional
2281    }
2282
2283    pub fn all_records(self) -> impl Iterator<Item = DnsRecordBox> {
2284        self.answers
2285            .into_iter()
2286            .chain(self.authorities)
2287            .chain(self.additional)
2288    }
2289
2290    pub fn num_additionals(&self) -> u16 {
2291        self.num_additionals
2292    }
2293
2294    pub fn num_authorities(&self) -> u16 {
2295        self.num_authorities
2296    }
2297
2298    pub fn num_questions(&self) -> u16 {
2299        self.num_questions
2300    }
2301
2302    pub const fn is_query(&self) -> bool {
2303        (self.flags & FLAGS_QR_MASK) == FLAGS_QR_QUERY
2304    }
2305
2306    pub const fn is_response(&self) -> bool {
2307        (self.flags & FLAGS_QR_MASK) == FLAGS_QR_RESPONSE
2308    }
2309
2310    fn read_header(&mut self) -> Result<()> {
2311        if self.data.len() < MSG_HEADER_LEN {
2312            return Err(e_fmt!(
2313                "DNS incoming: header is too short: {} bytes",
2314                self.data.len()
2315            ));
2316        }
2317
2318        let data = &self.data[0..];
2319        self.id = u16_from_be_slice(&data[..2]);
2320        self.flags = u16_from_be_slice(&data[2..4]);
2321        self.num_questions = u16_from_be_slice(&data[4..6]);
2322        self.num_answers = u16_from_be_slice(&data[6..8]);
2323        self.num_authorities = u16_from_be_slice(&data[8..10]);
2324        self.num_additionals = u16_from_be_slice(&data[10..12]);
2325
2326        self.offset = MSG_HEADER_LEN;
2327
2328        trace!(
2329            "read_header: id {}, {} questions {} answers {} authorities {} additionals",
2330            self.id,
2331            self.num_questions,
2332            self.num_answers,
2333            self.num_authorities,
2334            self.num_additionals
2335        );
2336        Ok(())
2337    }
2338
2339    fn read_questions(&mut self) -> Result<()> {
2340        trace!("read_questions: {}", &self.num_questions);
2341        for i in 0..self.num_questions {
2342            let name = self.read_name()?;
2343
2344            let data = &self.data[self.offset..];
2345            if data.len() < 4 {
2346                return Err(Error::Msg(format!(
2347                    "DNS incoming: question idx {} too short: {}",
2348                    i,
2349                    data.len()
2350                )));
2351            }
2352            let ty = u16_from_be_slice(&data[..2]);
2353            let class = u16_from_be_slice(&data[2..4]);
2354            self.offset += 4;
2355
2356            let Some(rr_type) = RRType::from_u16(ty) else {
2357                // Skip unsupported question types instead of failing the whole message.
2358                debug!("DNS incoming: skipping question idx {i} qtype unknown: {ty}");
2359                continue;
2360            };
2361
2362            self.questions.push(DnsQuestion {
2363                entry: DnsEntry::new(name, rr_type, class),
2364            });
2365        }
2366        Ok(())
2367    }
2368
2369    fn read_answers(&mut self) -> Result<()> {
2370        self.answers = self.read_rr_records(self.num_answers)?;
2371        Ok(())
2372    }
2373
2374    fn read_authorities(&mut self) -> Result<()> {
2375        self.authorities = self.read_rr_records(self.num_authorities)?;
2376        Ok(())
2377    }
2378
2379    fn read_additional(&mut self) -> Result<()> {
2380        self.additional = self.read_rr_records(self.num_additionals)?;
2381        Ok(())
2382    }
2383
2384    /// Decodes a sequence of RR records (in answers, authorities and additionals).
2385    fn read_rr_records(&mut self, count: u16) -> Result<Vec<DnsRecordBox>> {
2386        trace!("read_rr_records: {}", count);
2387        let mut rr_records = Vec::new();
2388
2389        // RFC 1035: https://datatracker.ietf.org/doc/html/rfc1035#section-3.2.1
2390        //
2391        // All RRs have the same top level format shown below:
2392        //                               1  1  1  1  1  1
2393        // 0  1  2  3  4  5  6  7  8  9  0  1  2  3  4  5
2394        // +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+
2395        // |                                               |
2396        // /                                               /
2397        // /                      NAME                     /
2398        // |                                               |
2399        // +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+
2400        // |                      TYPE                     |
2401        // +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+
2402        // |                     CLASS                     |
2403        // +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+
2404        // |                      TTL                      |
2405        // |                                               |
2406        // +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+
2407        // |                   RDLENGTH                    |
2408        // +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--|
2409        // /                     RDATA                     /
2410        // /                                               /
2411        // +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+
2412
2413        // Muse have at least TYPE, CLASS, TTL, RDLENGTH fields: 10 bytes.
2414        const RR_HEADER_REMAIN: usize = 10;
2415
2416        for _ in 0..count {
2417            let name = self.read_name()?;
2418            let slice = &self.data[self.offset..];
2419
2420            if slice.len() < RR_HEADER_REMAIN {
2421                return Err(Error::Msg(format!(
2422                    "read_others: RR '{}' is too short after name: {} bytes",
2423                    &name,
2424                    slice.len()
2425                )));
2426            }
2427
2428            let ty = u16_from_be_slice(&slice[..2]);
2429            let class = u16_from_be_slice(&slice[2..4]);
2430            let mut ttl = u32_from_be_slice(&slice[4..8]);
2431            if ttl == 0 && self.is_response() {
2432                // RFC 6762 section 10.1:
2433                // "...Queriers receiving a Multicast DNS response with a TTL of zero SHOULD
2434                // NOT immediately delete the record from the cache, but instead record
2435                // a TTL of 1 and then delete the record one second later."
2436                // See https://datatracker.ietf.org/doc/html/rfc6762#section-10.1
2437
2438                ttl = 1;
2439            }
2440            let rdata_len = u16_from_be_slice(&slice[8..10]) as usize;
2441            self.offset += RR_HEADER_REMAIN;
2442            let next_offset = self.offset + rdata_len;
2443
2444            // Sanity check for RDATA length.
2445            if next_offset > self.data.len() {
2446                return Err(Error::Msg(format!(
2447                    "RR {name} RDATA length {rdata_len} is invalid: remain data len: {}",
2448                    self.data.len() - self.offset
2449                )));
2450            }
2451
2452            // Decode the RDATA based on the record type. A single record with
2453            // malformed RDATA must not discard the whole message: skip just
2454            // that record and resume at the next one using RDLENGTH.
2455            match self.read_rdata(ty, class, ttl, rdata_len, &name) {
2456                Ok(Some(record)) => {
2457                    if self.offset == next_offset {
2458                        trace!("read_rr_records: {:?}", &record);
2459                        rr_records.push(record);
2460                    } else {
2461                        debug!(
2462                            "skipping record '{}' (type {}): RDATA ended at {}, expected {}",
2463                            &name, ty, self.offset, next_offset
2464                        );
2465                    }
2466                }
2467                Ok(None) => {
2468                    trace!("Unsupported DNS record type: {} name: {}", ty, &name);
2469                }
2470                Err(e) => {
2471                    debug!(
2472                        "skipping record '{}' (type {}) with invalid RDATA: {}",
2473                        &name, ty, e,
2474                    );
2475                }
2476            }
2477
2478            // Re-anchor to the record boundary defined by RDLENGTH, regardless
2479            // of how the RDATA decoded, so the next record is read from the
2480            // correct offset.
2481            self.offset = next_offset;
2482        }
2483
2484        Ok(rr_records)
2485    }
2486
2487    /// Decodes the RDATA of a single record whose header fields have already
2488    /// been read, returning `None` for record types we do not parse.
2489    ///
2490    /// On success the read cursor is left at the end of the RDATA; the caller
2491    /// verifies that against RDLENGTH. Errors are per-record: the caller skips
2492    /// the offending record and continues with the rest of the message.
2493    fn read_rdata(
2494        &mut self,
2495        ty: u16,
2496        class: u16,
2497        ttl: u32,
2498        rdata_len: usize,
2499        name: &str,
2500    ) -> Result<Option<DnsRecordBox>> {
2501        let rec: Option<DnsRecordBox> = match RRType::from_u16(ty) {
2502            None => None,
2503
2504            Some(rr_type) => match rr_type {
2505                RRType::CNAME | RRType::PTR => {
2506                    Some(DnsPointer::new(name, rr_type, class, ttl, self.read_name()?).boxed())
2507                }
2508                RRType::TXT => {
2509                    Some(DnsTxt::new(name, class, ttl, self.read_vec(rdata_len)?).boxed())
2510                }
2511                RRType::SRV => Some(
2512                    DnsSrv::new(
2513                        name,
2514                        class,
2515                        ttl,
2516                        self.read_u16()?,
2517                        self.read_u16()?,
2518                        self.read_u16()?,
2519                        self.read_name()?,
2520                    )
2521                    .boxed(),
2522                ),
2523                RRType::HINFO => Some(
2524                    DnsHostInfo::new(
2525                        name,
2526                        rr_type,
2527                        class,
2528                        ttl,
2529                        self.read_char_string()?,
2530                        self.read_char_string()?,
2531                    )
2532                    .boxed(),
2533                ),
2534                RRType::A => Some(
2535                    DnsAddress::new(
2536                        name,
2537                        rr_type,
2538                        class,
2539                        ttl,
2540                        self.read_ipv4()?.into(),
2541                        self.interface_id.clone(),
2542                    )
2543                    .boxed(),
2544                ),
2545                RRType::AAAA => Some(
2546                    DnsAddress::new(
2547                        name,
2548                        rr_type,
2549                        class,
2550                        ttl,
2551                        self.read_ipv6()?.into(),
2552                        self.interface_id.clone(),
2553                    )
2554                    .boxed(),
2555                ),
2556                RRType::NSEC => Some(
2557                    DnsNSec::new(
2558                        name,
2559                        class,
2560                        ttl,
2561                        self.read_name()?,
2562                        self.read_type_bitmap()?,
2563                    )
2564                    .boxed(),
2565                ),
2566                _ => None,
2567            },
2568        };
2569
2570        Ok(rec)
2571    }
2572
2573    fn read_char_string(&mut self) -> Result<String> {
2574        let Some(&length) = self.data.get(self.offset) else {
2575            return Err(e_fmt!(
2576                "read_char_string: no length byte at offset {}, data len {}",
2577                self.offset,
2578                self.data.len()
2579            ));
2580        };
2581        self.offset += 1;
2582        self.read_string(length as usize)
2583    }
2584
2585    fn read_u16(&mut self) -> Result<u16> {
2586        let slice = &self.data[self.offset..];
2587        if slice.len() < U16_SIZE {
2588            return Err(Error::Msg(format!(
2589                "read_u16: slice len is only {}",
2590                slice.len()
2591            )));
2592        }
2593        let num = u16_from_be_slice(&slice[..U16_SIZE]);
2594        self.offset += U16_SIZE;
2595        Ok(num)
2596    }
2597
2598    /// Reads the "Type Bit Map" block for a DNS NSEC record.
2599    fn read_type_bitmap(&mut self) -> Result<Vec<u8>> {
2600        // From RFC 6762: 6.1.  Negative Responses
2601        // https://datatracker.ietf.org/doc/html/rfc6762#section-6.1
2602        //   o The Type Bit Map block number is 0.
2603        //   o The Type Bit Map block length byte is a value in the range 1-32.
2604        //   o The Type Bit Map data is 1-32 bytes, as indicated by length
2605        //     byte.
2606
2607        // Sanity check: at least 2 bytes to read.
2608        if self.data.len() < self.offset + 2 {
2609            return Err(Error::Msg(format!(
2610                "DnsIncoming is too short: {} at NSEC Type Bit Map offset {}",
2611                self.data.len(),
2612                self.offset
2613            )));
2614        }
2615
2616        let block_num = self.data[self.offset];
2617        self.offset += 1;
2618        if block_num != 0 {
2619            return Err(Error::Msg(format!(
2620                "NSEC block number is not 0: {block_num}"
2621            )));
2622        }
2623
2624        let block_len = self.data[self.offset] as usize;
2625        if !(1..=32).contains(&block_len) {
2626            return Err(Error::Msg(format!(
2627                "NSEC block length must be in the range 1-32: {block_len}"
2628            )));
2629        }
2630        self.offset += 1;
2631
2632        let end = self.offset + block_len;
2633        if end > self.data.len() {
2634            return Err(Error::Msg(format!(
2635                "NSEC block overflow: {} over RData len {}",
2636                end,
2637                self.data.len()
2638            )));
2639        }
2640        let bitmap = self.data[self.offset..end].to_vec();
2641        self.offset += block_len;
2642
2643        Ok(bitmap)
2644    }
2645
2646    fn read_vec(&mut self, length: usize) -> Result<Vec<u8>> {
2647        if self.data.len() < self.offset + length {
2648            return Err(e_fmt!(
2649                "DNS Incoming: not enough data to read a chunk of data"
2650            ));
2651        }
2652
2653        let v = self.data[self.offset..self.offset + length].to_vec();
2654        self.offset += length;
2655        Ok(v)
2656    }
2657
2658    fn read_ipv4(&mut self) -> Result<Ipv4Addr> {
2659        if self.data.len() < self.offset + 4 {
2660            return Err(e_fmt!("DNS Incoming: not enough data to read an IPV4"));
2661        }
2662
2663        let bytes: [u8; 4] = self.data[self.offset..self.offset + 4]
2664            .try_into()
2665            .map_err(|_| e_fmt!("DNS incoming: Not enough bytes for reading an IPV4"))?;
2666        self.offset += bytes.len();
2667        Ok(Ipv4Addr::from(bytes))
2668    }
2669
2670    fn read_ipv6(&mut self) -> Result<Ipv6Addr> {
2671        if self.data.len() < self.offset + 16 {
2672            return Err(e_fmt!("DNS Incoming: not enough data to read an IPV6"));
2673        }
2674
2675        let bytes: [u8; 16] = self.data[self.offset..self.offset + 16]
2676            .try_into()
2677            .map_err(|_| e_fmt!("DNS incoming: Not enough bytes for reading an IPV6"))?;
2678        self.offset += bytes.len();
2679        Ok(Ipv6Addr::from(bytes))
2680    }
2681
2682    fn read_string(&mut self, length: usize) -> Result<String> {
2683        if self.data.len() < self.offset + length {
2684            return Err(e_fmt!("DNS Incoming: not enough data to read a string"));
2685        }
2686
2687        let s = str::from_utf8(&self.data[self.offset..self.offset + length])
2688            .map_err(|e| Error::Msg(e.to_string()))?;
2689        self.offset += length;
2690        Ok(s.to_string())
2691    }
2692
2693    /// Reads a domain name at the current location of `self.data`.
2694    ///
2695    /// See https://datatracker.ietf.org/doc/html/rfc1035#section-3.1 for
2696    /// domain name encoding.
2697    fn read_name(&mut self) -> Result<String> {
2698        let mut name = String::new();
2699        self.offset = self.read_labels(self.offset, &mut name)?;
2700        Ok(name)
2701    }
2702
2703    /// Appends the labels encoded at `offset` to `name`, and returns the offset
2704    /// just past that encoding: past the terminating zero byte, or past the
2705    /// compression pointer that ended the name.
2706    ///
2707    /// A name is a sequence of labels, where each label is a length byte
2708    /// followed by that many bytes. The name ends either with a zero length
2709    /// byte, or with a "compression pointer" (top 2 bits set) that redirects
2710    /// to a name written earlier in the same packet.
2711    ///
2712    /// For example, a packet where the question name `_http._tcp.local.` is
2713    /// written out in full at offset 12, and the answer name
2714    /// `myprinter._http._tcp.local.` at offset 40 reuses it via compression:
2715    ///
2716    /// ```text
2717    ///  offset:  12   13..17    18   19..22   23   24..28    29
2718    ///          +----+---------+----+--------+----+---------+----+
2719    ///  bytes:  | 05 | "_http" | 04 | "_tcp" | 05 | "local" | 00 |
2720    ///          +----+---------+----+--------+----+---------+----+
2721    ///            ^len           ^len          ^len           ^ zero byte: end of name
2722    ///
2723    ///  offset:  40    41..49     50   51
2724    ///          +----+-------------+----+----+
2725    ///  bytes:  | 09 | "myprinter" | C0 | 0C |
2726    ///          +----+-------------+----+----+
2727    ///            ^len               ^ pointer: 0xC00C ^ 0xC000 = 12, jump back to offset 12
2728    /// ```
2729    ///
2730    /// Takes `&self` so that following a pointer cannot move the read cursor.
2731    fn read_labels(&self, mut offset: usize, name: &mut String) -> Result<usize> {
2732        let data = &self.data[..];
2733
2734        // From RFC1035:
2735        // "...Domain names in messages are expressed in terms of a sequence of labels.
2736        // Each label is represented as a one octet length field followed by that
2737        // number of octets."
2738        //
2739        // "...The compression scheme allows a domain name in a message to be
2740        // represented as either:
2741        // - a sequence of labels ending in a zero octet
2742        // - a pointer
2743        // - a sequence of labels ending with a pointer"
2744        loop {
2745            if offset >= data.len() {
2746                return Err(Error::Msg(format!(
2747                    "read_labels: offset: {} data len {}",
2748                    offset,
2749                    data.len(),
2750                )));
2751            }
2752            let length = data[offset];
2753
2754            // From RFC1035:
2755            // "...a domain name is terminated by a length byte of zero."
2756            if length == 0 {
2757                return Ok(offset + 1); // The end of the name.
2758            }
2759
2760            // Check the first 2 bits for possible "Message compression".
2761            match length & 0xC0 {
2762                0x00 => {
2763                    // regular utf8 string with length
2764                    offset += 1;
2765                    let ending = offset + length as usize;
2766
2767                    // Never read beyond the whole data length.
2768                    if ending > data.len() {
2769                        return Err(Error::Msg(format!(
2770                            "read_labels: ending {} exceeds data length {}",
2771                            ending,
2772                            data.len()
2773                        )));
2774                    }
2775
2776                    let label = str::from_utf8(&data[offset..ending])
2777                        .map_err(|e| Error::Msg(format!("read_labels: from_utf8: {e}")))?;
2778
2779                    // `MAX_NAME_BYTES` bounds a possible loop where pointer targets a label that
2780                    // is already part of the current name. For example:
2781                    //
2782                    //  offset:  12   13..17    18   19
2783                    //          +----+---------+----+----+
2784                    //  bytes:  | 05 | "_http" | C0 | 0C |
2785                    //          +----+---------+----+----+
2786                    //            ^len           ^pointer targets offset 12.
2787                    if name.len() + label.len() + 1 > MAX_NAME_BYTES {
2788                        return Err(Error::Msg(format!(
2789                            "read_labels: name exceeds {MAX_NAME_BYTES} bytes: {name}"
2790                        )));
2791                    }
2792
2793                    *name += label;
2794                    *name += ".";
2795                    offset = ending;
2796                }
2797                0xC0 => {
2798                    // Message compression: a pointer marks the end of a domain name.
2799                    self.follow_pointer(offset, name)?;
2800                    return Ok(offset + U16_SIZE);
2801                }
2802                _ => {
2803                    return Err(Error::Msg(format!(
2804                        "Bad name with invalid length: 0x{:x} offset {}, data (so far): {:x?}",
2805                        length,
2806                        offset,
2807                        &data[..offset]
2808                    )));
2809                }
2810            };
2811        }
2812    }
2813
2814    /// Follows the compression pointer at offset `at`, appending the labels it
2815    /// names to `name`.
2816    ///
2817    /// See https://datatracker.ietf.org/doc/html/rfc1035#section-4.1.4 for
2818    /// message compression.
2819    fn follow_pointer(&self, at: usize, name: &mut String) -> Result<()> {
2820        let data = &self.data[..];
2821        let mut pointer_at = at;
2822
2823        // Resolve a run of pointers that target other pointers, so that the
2824        // recursive call below always lands on a label or on the end of a name.
2825        let target = loop {
2826            let slice = &data[pointer_at..];
2827            if slice.len() < U16_SIZE {
2828                return Err(Error::Msg(format!(
2829                    "follow_pointer: u16 slice len is only {}",
2830                    slice.len()
2831                )));
2832            }
2833            let target = (u16_from_be_slice(slice) ^ 0xC000) as usize;
2834
2835            // RFC1035 section 4.1.4 compresses a name into "a pointer to a prior
2836            // occurrence", so a pointer always points strictly backwards.
2837            if target >= pointer_at {
2838                return Err(Error::Msg(format!(
2839                    "Invalid name compression: pointer {target} at offset {pointer_at} must point backwards"
2840                )));
2841            }
2842
2843            if data[target] & 0xC0 != 0xC0 {
2844                break target;
2845            }
2846
2847            // The target is itself a pointer, so follow it.
2848            pointer_at = target;
2849        };
2850
2851        self.read_labels(target, name)?;
2852        Ok(())
2853    }
2854}
2855
2856const fn u16_from_be_slice(bytes: &[u8]) -> u16 {
2857    let u8_array: [u8; 2] = [bytes[0], bytes[1]];
2858    u16::from_be_bytes(u8_array)
2859}
2860
2861const fn u32_from_be_slice(s: &[u8]) -> u32 {
2862    let u8_array: [u8; 4] = [s[0], s[1], s[2], s[3]];
2863    u32::from_be_bytes(u8_array)
2864}
2865
2866/// Returns the time at which this record will have expired by a certain
2867/// percentage of its TTL.
2868fn get_expiration_time(created: Instant, ttl: u32, percent: u32) -> Instant {
2869    // 'ttl' is in seconds, hence:
2870    // ttl * 1000 * (percent / 100) => ttl * percent * 10
2871    created + Duration::from_millis(ttl as u64 * percent as u64 * 10)
2872}
2873
2874#[cfg(test)]
2875mod tests {
2876    use super::{
2877        u16_from_be_slice, DnsAddress, DnsHostInfo, DnsIncoming, DnsOutPacket, DnsOutgoing,
2878        DnsPointer, DnsTxt, RRType, CLASS_CACHE_FLUSH, CLASS_IN, FLAGS_QR_QUERY, FLAGS_QR_RESPONSE,
2879        FLAGS_TC, MAX_PKT_ABSOLUTE_IPV6, MAX_PKT_DEFAULT, MSG_HEADER_LEN,
2880    };
2881    use crate::InterfaceId;
2882    use std::collections::HashMap;
2883    use std::net::{IpAddr, Ipv4Addr};
2884
2885    /// The `is_ipv4` argument of `to_packets`. IPv6 has the smaller of the two
2886    /// absolute ceilings, so it is the stricter one to encode for.
2887    const IPV6: bool = false;
2888
2889    /// Found by fuzzing the packet parser.
2890    ///
2891    /// An HINFO record with RDLENGTH 0 placed at the very end of a message left
2892    /// `read_char_string` with no length octet to read, and it indexed one byte
2893    /// past the packet.
2894    #[test]
2895    fn test_hinfo_char_string_at_end_of_packet() {
2896        let mut data = Vec::new();
2897
2898        // Header: one authority record, and a query (so the TTL is not rewritten).
2899        data.extend_from_slice(&0x0087u16.to_be_bytes()); // id
2900        data.extend_from_slice(&0x0084u16.to_be_bytes()); // flags: a query
2901        data.extend_from_slice(&0u16.to_be_bytes()); // 0 questions
2902        data.extend_from_slice(&0u16.to_be_bytes()); // 0 answers
2903        data.extend_from_slice(&1u16.to_be_bytes()); // 1 authorities
2904        data.extend_from_slice(&0u16.to_be_bytes()); // 0 additionals
2905
2906        data.push(0); // name: root
2907        data.extend_from_slice(&(RRType::HINFO as u16).to_be_bytes());
2908        data.extend_from_slice(&CLASS_IN.to_be_bytes());
2909        data.extend_from_slice(&0u32.to_be_bytes()); // ttl
2910
2911        // RDLENGTH is 0, so the record โ€” and the message โ€” end here, leaving
2912        // nothing for HINFO's two <character-string> fields.
2913        data.extend_from_slice(&0u16.to_be_bytes()); // rdlength
2914
2915        assert_eq!(data.len(), 23);
2916
2917        let parsed = DnsIncoming::new(data, test_interface_id())
2918            .expect("a truncated HINFO must be skipped, not fail the packet");
2919
2920        // The record is dropped, and nothing is left behind.
2921        assert_eq!(parsed.authorities().len(), 0);
2922    }
2923
2924    #[test]
2925    fn test_dns_outgoing_serialization_empty() {
2926        let out = DnsOutgoing::new(0);
2927        let packets = out.to_packets(MAX_PKT_DEFAULT, IPV6);
2928        assert_eq!(packets.len(), 1);
2929        assert_eq!(packets[0].as_bytes(), &[0; 12]);
2930        let expected_names = HashMap::new();
2931        assert_eq!(&packets[0].names, &expected_names);
2932    }
2933
2934    #[test]
2935    fn test_dns_outgoing_serialization_question() {
2936        let mut out = DnsOutgoing::new(0);
2937        out.add_question("123.test", RRType::A);
2938        let packets = out.to_packets(MAX_PKT_DEFAULT, IPV6);
2939        assert_eq!(packets.len(), 1);
2940        assert_eq!(
2941            packets[0].as_bytes(),
2942            &[
2943                0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, // Header
2944                // Payload
2945                3, 49, 50, 51, 4, 116, 101, 115, 116, 0, 0, 1, 0, 1,
2946            ]
2947        );
2948        let mut expected_names = HashMap::new();
2949        expected_names.insert("123.test".to_string(), 12);
2950        expected_names.insert("test".to_string(), 16);
2951        assert_eq!(&packets[0].names, &expected_names);
2952    }
2953
2954    #[test]
2955    fn test_dns_outgoing_serialization_question_with_authority() {
2956        let mut out = DnsOutgoing::new(0);
2957        out.add_question("123.test", RRType::ANY);
2958        out.add_authority(Box::new(DnsTxt::new(
2959            "124.test",
2960            CLASS_IN,
2961            0x00112233,
2962            b"help".to_vec(),
2963        )));
2964        out.add_authority(Box::new(DnsHostInfo::new(
2965            "124.test",
2966            RRType::CNAME,
2967            CLASS_IN,
2968            0x00112233,
2969            "arm".to_string(),
2970            "linux".to_string(),
2971        )));
2972        let packets = out.to_packets(MAX_PKT_DEFAULT, IPV6);
2973        assert_eq!(packets.len(), 1);
2974        assert_eq!(
2975            packets[0].as_bytes(),
2976            &[
2977                0, 0, 0, 0, 0, 1, 0, 0, 0, 2, 0, 0, // Header
2978                // Payload
2979                3, 49, 50, 51, 4, 116, 101, 115, 116, 0, 0, 255, 0, 1, 3, 49, 50, 52, 192, 16, 0,
2980                16, 0, 1, 0, 17, 34, 51, 0, 4, 104, 101, 108, 112, 192, 26, 0, 5, 0, 1, 0, 17, 34,
2981                51, 0, 8, 97, 114, 109, 108, 105, 110, 117, 120,
2982            ]
2983        );
2984        let mut expected_names = HashMap::new();
2985        expected_names.insert("123.test".to_string(), 12);
2986        expected_names.insert("test".to_string(), 16);
2987        expected_names.insert("124.test".to_string(), 26);
2988        assert_eq!(&packets[0].names, &expected_names);
2989    }
2990
2991    #[test]
2992    fn test_dns_outgoing_serialization_additional_answer() {
2993        let mut out = DnsOutgoing::new(0);
2994        out.add_additional_answer(DnsAddress::new(
2995            "test.local",
2996            RRType::A,
2997            CLASS_IN | CLASS_CACHE_FLUSH,
2998            0xdead_beef,
2999            IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1)),
3000            InterfaceId::default(),
3001        ));
3002        let packets = out.to_packets(MAX_PKT_DEFAULT, IPV6);
3003        assert_eq!(packets.len(), 1);
3004        assert_eq!(
3005            packets[0].as_bytes(),
3006            &[
3007                0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, // Header
3008                // Payload
3009                4, 116, 101, 115, 116, 5, 108, 111, 99, 97, 108, 0, 0, 1, 128, 1, 222, 173, 190,
3010                239, 0, 4, 127, 0, 0, 1,
3011            ]
3012        );
3013        let mut expected_names = HashMap::new();
3014        expected_names.insert("test.local".to_string(), 12);
3015        expected_names.insert("local".to_string(), 17);
3016        assert_eq!(&packets[0].names, &expected_names);
3017    }
3018
3019    #[test]
3020    fn test_dns_outgoing_serialization_answer_at_time() {
3021        let mut out = DnsOutgoing::new(0);
3022        out.add_answer_record(DnsPointer::new(
3023            "test",
3024            RRType::PTR,
3025            CLASS_IN,
3026            0xaaaa5555,
3027            "test-service".to_string(),
3028        ));
3029        let packets = out.to_packets(MAX_PKT_DEFAULT, IPV6);
3030        assert_eq!(packets.len(), 1);
3031        assert_eq!(
3032            packets[0].as_bytes(),
3033            &[
3034                0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0, // Header
3035                // Payload
3036                4, 116, 101, 115, 116, 0, 0, 12, 0, 1, 170, 170, 85, 85, 0, 14, 12, 116, 101, 115,
3037                116, 45, 115, 101, 114, 118, 105, 99, 101, 0,
3038            ]
3039        );
3040
3041        let mut out = DnsOutgoing::new(0);
3042        out.add_answer_record(DnsPointer::new(
3043            "test",
3044            RRType::CNAME,
3045            CLASS_IN,
3046            0xaaaa5555,
3047            "test-service.local".to_string(),
3048        ));
3049        out.add_answer_record(DnsPointer::new(
3050            "test",
3051            RRType::AAAA,
3052            CLASS_IN,
3053            0xffffffff,
3054            "test-service.local".to_string(),
3055        ));
3056        let packets = out.to_packets(MAX_PKT_DEFAULT, IPV6);
3057        assert_eq!(packets.len(), 1);
3058        assert_eq!(
3059            packets[0].as_bytes(),
3060            &[
3061                0, 0, 0, 0, 0, 0, 0, 2, 0, 0, 0, 0, // Header
3062                // Payload
3063                4, 116, 101, 115, 116, 0, 0, 5, 0, 1, 170, 170, 85, 85, 0, 20, 12, 116, 101, 115,
3064                116, 45, 115, 101, 114, 118, 105, 99, 101, 5, 108, 111, 99, 97, 108, 0, 192, 12, 0,
3065                28, 0, 1, 255, 255, 255, 255, 0, 2, 192, 28,
3066            ]
3067        );
3068        let mut expected_names = HashMap::new();
3069        expected_names.insert("test".to_string(), 12);
3070        expected_names.insert("test-service.local".to_string(), 28);
3071        expected_names.insert("local".to_string(), 41);
3072        assert_eq!(&packets[0].names, &expected_names);
3073    }
3074
3075    /// A question whose name has a label longer than 63 bytes cannot be
3076    /// encoded. It must be skipped, not panic. (Note the question count in the
3077    /// header must reflect the questions actually written.)
3078    #[test]
3079    fn test_dns_outgoing_question_label_too_long() {
3080        let long_label = "a".repeat(64);
3081        let mut out = DnsOutgoing::new(0);
3082        out.add_question(&format!("{long_label}.local"), RRType::PTR);
3083        out.add_question("123.test", RRType::A);
3084
3085        let packets = out.to_packets(MAX_PKT_DEFAULT, IPV6);
3086        assert_eq!(packets.len(), 1);
3087        assert_eq!(
3088            packets[0].as_bytes(),
3089            &[
3090                0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, // Header: 1 question
3091                // Payload: only "123.test" made it in.
3092                3, 49, 50, 51, 4, 116, 101, 115, 116, 0, 0, 1, 0, 1,
3093            ]
3094        );
3095
3096        // The rolled back name must not leave a stale compression offset behind.
3097        let mut expected_names = HashMap::new();
3098        expected_names.insert("123.test".to_string(), 12);
3099        expected_names.insert("test".to_string(), 16);
3100        assert_eq!(&packets[0].names, &expected_names);
3101    }
3102
3103    /// A record whose rdata carries an unencodable name (here a PTR alias) is
3104    /// dropped as a whole, leaving the rest of the packet intact.
3105    #[test]
3106    fn test_dns_outgoing_record_label_too_long() {
3107        let long_label = "a".repeat(64);
3108        let mut out = DnsOutgoing::new(0);
3109        out.add_answer_record(DnsPointer::new(
3110            "_test._tcp.local.",
3111            RRType::PTR,
3112            CLASS_IN,
3113            0,
3114            format!("{long_label}._test._tcp.local."),
3115        ));
3116        out.add_answer_record(DnsPointer::new(
3117            "_test._tcp.local.",
3118            RRType::PTR,
3119            CLASS_IN,
3120            0,
3121            "ok._test._tcp.local.".to_string(),
3122        ));
3123
3124        let packets = out.to_packets(MAX_PKT_DEFAULT, IPV6);
3125        assert_eq!(packets.len(), 1);
3126
3127        // Header answer count is 1: the first answer was dropped.
3128        assert_eq!(&packets[0].as_bytes()[6..8], &[0, 1]);
3129
3130        // Re-parsing must succeed and yield only the good answer.
3131        let incoming = DnsIncoming::new(
3132            packets[0].as_bytes().to_vec(),
3133            InterfaceId {
3134                name: "test".to_string(),
3135                index: 1,
3136            },
3137        )
3138        .unwrap();
3139        assert_eq!(incoming.answers().len(), 1);
3140    }
3141
3142    /// A name learned from the network can hold a label that ends with a
3143    /// backslash, which escapes the following label separator. Unescaping such
3144    /// a name on the way out merges two 63-byte labels into a 127-byte one.
3145    /// This used to panic the daemon thread. See issue #483.
3146    #[test]
3147    fn test_incoming_name_with_merged_labels_does_not_panic() {
3148        // A query with one question: "aa..a\" + "bb..b", 63 bytes each.
3149        let mut data: Vec<u8> = vec![0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0];
3150        data.push(63);
3151        data.extend(vec![b'a'; 62]);
3152        data.push(b'\\');
3153        data.push(63);
3154        data.extend(vec![b'b'; 63]);
3155        data.push(0);
3156        data.extend([0, 12, 0, 1]); // PTR, IN
3157
3158        let incoming = DnsIncoming::new(
3159            data,
3160            InterfaceId {
3161                name: "test".to_string(),
3162                index: 1,
3163            },
3164        )
3165        .unwrap();
3166        let name = incoming.questions()[0].entry.name.clone();
3167
3168        // The two labels merged: the trailing backslash escaped the separator.
3169        assert!(name.starts_with("aaa"));
3170        assert!(name.contains("\\.bbb"));
3171
3172        // Re-emitting it must drop the question rather than panic.
3173        let mut out = DnsOutgoing::new(0);
3174        out.add_question(&name, RRType::PTR);
3175        let packets = out.to_packets(MAX_PKT_DEFAULT, IPV6);
3176        assert_eq!(packets.len(), 1);
3177        assert_eq!(packets[0].as_bytes(), &[0; MSG_HEADER_LEN]);
3178    }
3179
3180    /// A pointer that points into the name currently being read is a loop:
3181    /// following it re-reads the same labels and arrives at the same pointer
3182    /// again. `read_name` must reject such a name instead of hanging.
3183    #[test]
3184    fn test_read_name_pointer_loop_is_rejected() {
3185        // A response with one PTR record. Its name starts at offset 12 and is
3186        // encoded as: label "local", label "_x", then a pointer back to 12,
3187        // i.e. to the "local" label of this very name.
3188        let mut data: Vec<u8> = vec![0, 0, 0x84, 0, 0, 0, 0, 1, 0, 0, 0, 0];
3189        data.extend_from_slice(&[5, b'l', b'o', b'c', b'a', b'l']); // offset 12
3190        data.extend_from_slice(&[2, b'_', b'x']); // offset 18
3191        data.extend_from_slice(&[0xC0, 12]); // offset 21: pointer to 12
3192        data.extend_from_slice(&[0, 12, 0, 1]); // PTR, IN
3193        data.extend_from_slice(&[0, 0, 0, 120]); // TTL
3194        data.extend_from_slice(&[0, 2]); // RDLENGTH
3195        data.extend_from_slice(&[0xC0, 12]); // RDATA: pointer to 12
3196
3197        assert!(DnsIncoming::new(data, test_interface_id()).is_err());
3198    }
3199
3200    /// A legal name that follows a pointer backwards and then meets a second
3201    /// pointer whose target sits *after* the start of the name being read, yet
3202    /// still strictly *before* that second pointer's own position.
3203    ///
3204    /// Such a message probably never appears in reality, but it still has to parse.
3205    /// Reading the answer's name walks: 700 -> 640 -> 62-byte label -> 703 ->
3206    /// 702 -> zero byte, name complete.
3207    #[test]
3208    fn test_read_name_pointer_after_backward_jump() {
3209        /// Appends a question: one label of `label_len` 'a' bytes, PTR, IN.
3210        fn push_question(data: &mut Vec<u8>, label_len: usize) {
3211            data.push(label_len as u8);
3212            data.extend(vec![b'a'; label_len]);
3213            data.push(0); // end of the name
3214            data.extend_from_slice(&[0, 12]); // QTYPE: PTR
3215            data.extend_from_slice(&[0, 1]); // QCLASS: IN
3216        }
3217
3218        let mut data: Vec<u8> = vec![
3219            0, 0, // ID
3220            0, 0, // flags: a query
3221            0, 11, // 11 questions
3222            0, 1, // 1 answer
3223            0, 0, 0, 0, // no authorities, no additionals
3224        ];
3225
3226        // Questions #1 to #10, 66 bytes each: 12 + 660 = 672.
3227        for _ in 0..10 {
3228            push_question(&mut data, 60);
3229        }
3230        assert_eq!(data.len(), 672);
3231
3232        // Question #11, 28 bytes, so that the answer record starts at 700.
3233        push_question(&mut data, 22);
3234        assert_eq!(data.len(), 700);
3235
3236        // Plant the label length inside question #10's label.
3237        data[640] = 62;
3238
3239        // The answer record.
3240        data.extend_from_slice(&[0xC2, 0x80]); // 700: name: pointer to 640
3241        data.extend_from_slice(&[0x00, 0xC2]); // 702: TYPE, unknown type 194
3242        data.extend_from_slice(&[0xBE, 0x01]); // 704: CLASS. 703..705 is a pointer to 702
3243        data.extend_from_slice(&[0, 0, 0, 120]); // TTL
3244        data.extend_from_slice(&[0, 0]); // RDLENGTH: no RDATA
3245
3246        // Both pointers point backwards from where they are.
3247        assert_eq!(u16_from_be_slice(&data[700..702]) ^ 0xC000, 640);
3248        assert_eq!(u16_from_be_slice(&data[703..705]) ^ 0xC000, 702);
3249
3250        let incoming = DnsIncoming::new(data, test_interface_id())
3251            .expect("a name whose pointers all point backwards must parse");
3252        assert_eq!(incoming.questions().len(), 11);
3253
3254        // The answer's type is unknown to us, so the record itself is skipped.
3255        assert_eq!(incoming.answers().len(), 0);
3256    }
3257
3258    /// Two pointers at offsets 23 and 25 that target each other (23 -> 25 ->
3259    /// 23). Both sit below offset 27, where the name starts.
3260    ///
3261    /// `follow_pointer` requires each target to be strictly below the
3262    /// pointer's *own* position. A cycle always contains at least one
3263    /// non-backward hop, so this rule breaks every cycle.
3264    #[test]
3265    fn test_read_name_mutual_pointers_are_rejected() {
3266        let mut data: Vec<u8> = vec![0, 0, 0x84, 0, 0, 0, 0, 2, 0, 0, 0, 0];
3267
3268        // Answer #1: the root name, then an unknown type, so its RDATA is skipped.
3269        data.push(0); // 12: the root name
3270        data.extend_from_slice(&[0x00, 0xC2]); // 13: TYPE: unknown type 194
3271        data.extend_from_slice(&[0x00, 0x01]); // 15: CLASS: IN
3272        data.extend_from_slice(&[0, 0, 0, 120]); // 17: TTL
3273        data.extend_from_slice(&[0x00, 0x04]); // 21: RDLENGTH
3274        data.extend_from_slice(&[0xC0, 25]); // 23: RDATA: pointer to 25
3275        data.extend_from_slice(&[0xC0, 23]); // 25: RDATA: pointer to 23
3276        assert_eq!(data.len(), 27);
3277
3278        // Answer #2, whose name points into that RDATA.
3279        data.extend_from_slice(&[0xC0, 23]); // 27: name: pointer to 23
3280        data.extend_from_slice(&[0x00, 0xC2, 0x00, 0x01]); // TYPE, CLASS
3281        data.extend_from_slice(&[0, 0, 0, 120]); // TTL
3282        data.extend_from_slice(&[0, 0]); // RDLENGTH: no RDATA
3283
3284        // Every pointer targets an offset below the start of the name at 27.
3285        assert_eq!(u16_from_be_slice(&data[27..29]) ^ 0xC000, 23);
3286        assert_eq!(u16_from_be_slice(&data[23..25]) ^ 0xC000, 25);
3287        assert_eq!(u16_from_be_slice(&data[25..27]) ^ 0xC000, 23);
3288
3289        assert!(DnsIncoming::new(data, test_interface_id()).is_err());
3290    }
3291
3292    /// A label whose read carries the cursor onto a pointer that jumps back to
3293    /// that same label. Every pointer here points backwards from its own
3294    /// position, so no comparison of offsets rejects it: the cycle is broken
3295    /// only by the name growing past [`MAX_NAME_BYTES`].
3296    #[test]
3297    fn test_read_name_label_cycle_is_rejected() {
3298        let mut data: Vec<u8> = vec![0, 0, 0x84, 0, 0, 0, 0, 2, 0, 0, 0, 0];
3299
3300        // Answer #1, again an unknown type so that its RDATA is skipped.
3301        data.push(0); // 12: the root name
3302        data.extend_from_slice(&[0x00, 0xC2]); // 13: TYPE: unknown type 194
3303        data.extend_from_slice(&[0x00, 0x01]); // 15: CLASS: IN
3304        data.extend_from_slice(&[0, 0, 0, 120]); // 17: TTL
3305        data.extend_from_slice(&[0x00, 0x07]); // 21: RDLENGTH
3306        data.push(0x04); // 23: RDATA: a label of 4 bytes, ending at 28
3307        data.extend_from_slice(b"aaaa"); // 24
3308        data.extend_from_slice(&[0xC0, 23]); // 28: RDATA: pointer to 23
3309        assert_eq!(data.len(), 30);
3310
3311        // Answer #2, whose name enters the cycle.
3312        data.extend_from_slice(&[0xC0, 23]); // 30: name: pointer to 23
3313        data.extend_from_slice(&[0x00, 0xC2, 0x00, 0x01]); // TYPE, CLASS
3314        data.extend_from_slice(&[0, 0, 0, 120]); // TTL
3315        data.extend_from_slice(&[0, 0]); // RDLENGTH: no RDATA
3316
3317        // Reading the label at 23 leaves the cursor on the pointer at 28, which
3318        // points backwards from 28 and lands back on the label.
3319        assert_eq!(u16_from_be_slice(&data[28..30]) ^ 0xC000, 23);
3320        assert_eq!(u16_from_be_slice(&data[30..32]) ^ 0xC000, 23);
3321
3322        assert!(DnsIncoming::new(data, test_interface_id()).is_err());
3323    }
3324
3325    /// A real `_miio._udp.local.` response captured behind an avahi mDNS
3326    /// reflector (see issue #468). It has 5 answers, one of which is an NSEC
3327    /// whose Next Domain Name is a compression pointer to its own offset (a
3328    /// self-reference, offset 121 -> 121). That one record is malformed, but
3329    /// the other four (PTR, A, SRV, TXT) are fine, and lenient parsers such as
3330    /// tcpdump decode the whole packet.
3331    ///
3332    /// The parser must skip only the malformed NSEC and keep the good records,
3333    /// rather than discarding the entire message.
3334    #[test]
3335    fn test_malformed_nsec_record_is_skipped() {
3336        let data: Vec<u8> = vec![
3337            0x00, 0x00, 0x84, 0x00, 0x00, 0x00, 0x00, 0x05, 0x00, 0x00, 0x00, 0x00, 0x05, 0x5f,
3338            0x6d, 0x69, 0x69, 0x6f, 0x04, 0x5f, 0x75, 0x64, 0x70, 0x05, 0x6c, 0x6f, 0x63, 0x61,
3339            0x6c, 0x00, 0x00, 0x0c, 0x00, 0x01, 0x00, 0x00, 0x00, 0x78, 0x00, 0x24, 0x21, 0x64,
3340            0x72, 0x65, 0x61, 0x6d, 0x65, 0x2d, 0x76, 0x61, 0x63, 0x75, 0x75, 0x6d, 0x2d, 0x70,
3341            0x32, 0x30, 0x32, 0x39, 0x5f, 0x6d, 0x69, 0x69, 0x6f, 0x34, 0x34, 0x37, 0x33, 0x30,
3342            0x35, 0x32, 0x34, 0x37, 0xc0, 0x0c, 0x21, 0x64, 0x72, 0x65, 0x61, 0x6d, 0x65, 0x2d,
3343            0x76, 0x61, 0x63, 0x75, 0x75, 0x6d, 0x2d, 0x70, 0x32, 0x30, 0x32, 0x39, 0x5f, 0x6d,
3344            0x69, 0x69, 0x6f, 0x34, 0x34, 0x37, 0x33, 0x30, 0x35, 0x32, 0x34, 0x37, 0x00, 0x00,
3345            0x2f, 0x80, 0x01, 0x00, 0x00, 0x00, 0x78, 0x00, 0x09, 0xc0, 0x79, 0x00, 0x05, 0x40,
3346            0x00, 0x00, 0x00, 0x00, 0xc0, 0x4c, 0x00, 0x01, 0x80, 0x01, 0x00, 0x00, 0x00, 0x78,
3347            0x00, 0x04, 0x0a, 0x2a, 0x02, 0x32, 0xc0, 0x28, 0x00, 0x21, 0x80, 0x01, 0x00, 0x00,
3348            0x00, 0x78, 0x00, 0x08, 0x00, 0x00, 0x00, 0x00, 0xd4, 0x31, 0xc0, 0x4c, 0xc0, 0x28,
3349            0x00, 0x10, 0x80, 0x01, 0x00, 0x00, 0x00, 0x78, 0x00, 0x0f, 0x0e, 0x70, 0x61, 0x74,
3350            0x68, 0x3d, 0x2f, 0x6d, 0x79, 0x64, 0x65, 0x76, 0x69, 0x63, 0x65,
3351        ];
3352
3353        // The offending record: the NSEC's Next Domain Name at offset 121 is a
3354        // pointer to offset 121 (itself).
3355        assert_eq!(u16_from_be_slice(&data[121..123]) ^ 0xC000, 121);
3356
3357        let incoming = DnsIncoming::new(data, test_interface_id())
3358            .expect("one malformed record must not fail the whole packet");
3359
3360        // Four of the five records survive; only the NSEC is dropped.
3361        assert_eq!(incoming.answers().len(), 4);
3362        assert!(
3363            !incoming
3364                .answers()
3365                .iter()
3366                .any(|r| r.get_type() == RRType::NSEC),
3367            "the malformed NSEC record must be skipped"
3368        );
3369    }
3370
3371    #[test]
3372    fn test_unknown_questions_preserve_supported_questions_and_answers() {
3373        // No enum variant is needed for future question types. Build the wire
3374        // questions explicitly so this also covers types the writer cannot emit.
3375        for unknown_type in [64u16, 65, 65400] {
3376            for unknown_index in 0..3 {
3377                let mut data = vec![0; 12];
3378                data[4..6].copy_from_slice(&3u16.to_be_bytes());
3379                data[6..8].copy_from_slice(&1u16.to_be_bytes());
3380                let mut known_types = [1u16, 28].iter().copied();
3381                for index in 0..3 {
3382                    data.extend_from_slice(b"\x05mixed\x05local\x00");
3383                    let ty = if index == unknown_index {
3384                        unknown_type
3385                    } else {
3386                        known_types.next().unwrap()
3387                    };
3388                    data.extend_from_slice(&ty.to_be_bytes());
3389                    data.extend_from_slice(&CLASS_IN.to_be_bytes());
3390                }
3391                // One known answer after the questions verifies parser alignment.
3392                data.extend_from_slice(
3393                    b"\xc0\x0c\x00\x01\x00\x01\x00\x00\x00\x78\x00\x04\xc0\x00\x02\x01",
3394                );
3395                let incoming = DnsIncoming::new(data, test_interface_id()).unwrap();
3396                let types: Vec<_> = incoming
3397                    .questions()
3398                    .iter()
3399                    .map(|q| q.entry.ty)
3400                    .filter(|ty| matches!(ty, RRType::A | RRType::AAAA))
3401                    .collect();
3402                assert_eq!(types, [RRType::A, RRType::AAAA]);
3403                if unknown_type == 65400 {
3404                    assert_eq!(incoming.questions().len(), 2);
3405                }
3406                assert_eq!(incoming.answers().len(), 1);
3407                assert_eq!(incoming.answers()[0].get_type(), RRType::A);
3408            }
3409        }
3410    }
3411
3412    #[test]
3413    fn test_unknown_question_still_requires_complete_name_type_and_class() {
3414        let mut data = vec![0; 12];
3415        data[4..6].copy_from_slice(&1u16.to_be_bytes());
3416        data.extend_from_slice(b"\x05mixed\x05local\x00");
3417        data.extend_from_slice(&65400u16.to_be_bytes());
3418        data.extend_from_slice(&CLASS_IN.to_be_bytes());
3419        assert!(DnsIncoming::new(data.clone(), test_interface_id())
3420            .unwrap()
3421            .questions()
3422            .is_empty());
3423        for missing in 1..=4 {
3424            assert!(
3425                DnsIncoming::new(data[..data.len() - missing].to_vec(), test_interface_id())
3426                    .is_err()
3427            );
3428        }
3429    }
3430
3431    fn test_interface_id() -> InterfaceId {
3432        InterfaceId {
3433            name: "test".to_string(),
3434            index: 1,
3435        }
3436    }
3437
3438    /// The "flags" field of a finished packet.
3439    fn packet_flags(packet: &DnsOutPacket) -> u16 {
3440        let bytes = packet.as_bytes();
3441        u16::from_be_bytes([bytes[2], bytes[3]])
3442    }
3443
3444    fn ptr_answer(index: usize) -> DnsPointer {
3445        DnsPointer::new(
3446            "_spill._tcp.local.",
3447            RRType::PTR,
3448            CLASS_IN,
3449            4500,
3450            format!("instance-{index:04}._spill._tcp.local."),
3451        )
3452    }
3453
3454    /// Re-parses each packet and returns the total number of answers found, which
3455    /// checks the header counts against what each packet actually holds.
3456    fn parsed_answer_count(packets: &[DnsOutPacket]) -> usize {
3457        packets
3458            .iter()
3459            .map(|packet: &DnsOutPacket| {
3460                let parsed = DnsIncoming::new(packet.as_bytes().to_vec(), test_interface_id())
3461                    .expect("each packet must parse on its own");
3462                assert!(
3463                    !parsed.answers().is_empty(),
3464                    "a spilled packet must not be empty"
3465                );
3466                parsed.answers().len()
3467            })
3468            .sum()
3469    }
3470
3471    /// A response too big for one packet spills into more packets. Every record
3472    /// must survive: before, records that did not fit were silently dropped.
3473    #[test]
3474    fn test_dns_outgoing_response_spills_into_packets() {
3475        const ANSWER_COUNT: usize = 100;
3476
3477        let mut out = DnsOutgoing::new(FLAGS_QR_RESPONSE);
3478        for i in 0..ANSWER_COUNT {
3479            out.add_answer_record(ptr_answer(i));
3480        }
3481
3482        let packets = out.to_packets(MAX_PKT_DEFAULT, IPV6);
3483        assert!(
3484            packets.len() > 1,
3485            "{} answers should not fit in one packet",
3486            ANSWER_COUNT
3487        );
3488
3489        for packet in &packets {
3490            assert!(
3491                packet.size() <= MAX_PKT_DEFAULT,
3492                "packet of {} bytes exceeds the limit",
3493                packet.size()
3494            );
3495
3496            // A multi-packet response is a series of independent responses: unlike
3497            // a query's known answers, it does not use the TC bit.
3498            assert_eq!(packet_flags(packet) & FLAGS_TC, 0);
3499        }
3500
3501        assert_eq!(parsed_answer_count(&packets), ANSWER_COUNT);
3502    }
3503
3504    /// RFC 6762 section 7.2: a querier sending known answers in more than one
3505    /// packet sets the TC bit in every packet but the last.
3506    #[test]
3507    fn test_dns_outgoing_query_truncation_bit() {
3508        let mut out = DnsOutgoing::new(FLAGS_QR_QUERY);
3509        out.add_question("_spill._tcp.local.", RRType::PTR);
3510        for i in 0..100 {
3511            out.add_answer_box(Box::new(ptr_answer(i)));
3512        }
3513
3514        let packets = out.to_packets(MAX_PKT_DEFAULT, IPV6);
3515        assert!(
3516            packets.len() > 1,
3517            "known answers should not fit in one packet"
3518        );
3519
3520        let (last, rest) = packets.split_last().expect("at least one packet");
3521        for packet in rest {
3522            assert_ne!(
3523                packet_flags(packet) & FLAGS_TC,
3524                0,
3525                "a packet with more known answers to follow must set TC"
3526            );
3527        }
3528        assert_eq!(
3529            packet_flags(last) & FLAGS_TC,
3530            0,
3531            "the last packet must not set TC"
3532        );
3533
3534        // The question goes in the first packet only, and no answer is lost.
3535        assert_eq!(packets[0].as_bytes()[4..6], 1u16.to_be_bytes());
3536        for packet in rest.iter().skip(1) {
3537            assert_eq!(packet.as_bytes()[4..6], [0, 0]);
3538        }
3539        assert_eq!(parsed_answer_count(&packets), 100);
3540    }
3541
3542    /// RFC 6762 section 17: a record too large for one MTU-sized packet is sent
3543    /// alone in an oversized packet, rather than dropped. It must be alone, since
3544    /// a fragmented packet "MUST NOT contain more than one resource record".
3545    #[test]
3546    fn test_dns_outgoing_oversized_record_sent_alone() {
3547        let mut out = DnsOutgoing::new(FLAGS_QR_RESPONSE);
3548        out.add_answer_record(ptr_answer(0));
3549        out.add_answer_record(DnsTxt::new(
3550            "big._spill._tcp.local.",
3551            CLASS_IN,
3552            4500,
3553            vec![b'x'; 2000],
3554        ));
3555        out.add_answer_record(ptr_answer(1));
3556
3557        let packets = out.to_packets(MAX_PKT_DEFAULT, IPV6);
3558        assert_eq!(packets.len(), 3, "the big record needs a packet to itself");
3559
3560        assert!(packets[0].size() <= MAX_PKT_DEFAULT);
3561        assert!(
3562            packets[1].size() > MAX_PKT_DEFAULT,
3563            "the oversized record must not be dropped"
3564        );
3565        // Still small enough that the send path will let it out.
3566        assert!(packets[1].size() <= MAX_PKT_ABSOLUTE_IPV6);
3567        assert!(packets[2].size() <= MAX_PKT_DEFAULT);
3568
3569        // One record per packet here, the middle one being the big TXT.
3570        let parsed = DnsIncoming::new(packets[1].as_bytes().to_vec(), test_interface_id()).unwrap();
3571        assert_eq!(parsed.answers().len(), 1);
3572        assert_eq!(parsed.answers()[0].get_name(), "big._spill._tcp.local.");
3573        assert_eq!(parsed_answer_count(&packets), 3);
3574    }
3575
3576    /// A record over the RFC 6762 section 17 ceiling could not go out on the wire
3577    /// even in a packet of its own, so it is dropped while its neighbors survive.
3578    #[test]
3579    fn test_dns_outgoing_record_over_absolute_ceiling_dropped() {
3580        let mut out = DnsOutgoing::new(FLAGS_QR_RESPONSE);
3581        out.add_answer_record(ptr_answer(0));
3582        out.add_answer_record(DnsTxt::new(
3583            "huge._spill._tcp.local.",
3584            CLASS_IN,
3585            4500,
3586            vec![b'x'; MAX_PKT_ABSOLUTE_IPV6],
3587        ));
3588        out.add_answer_record(ptr_answer(1));
3589
3590        let packets = out.to_packets(MAX_PKT_DEFAULT, IPV6);
3591        for packet in &packets {
3592            assert!(
3593                packet.size() <= MAX_PKT_ABSOLUTE_IPV6,
3594                "an unsendable packet must never be generated"
3595            );
3596        }
3597        assert_eq!(
3598            parsed_answer_count(&packets),
3599            2,
3600            "only the huge record is dropped"
3601        );
3602    }
3603
3604    /// Authorities and additionals spill too, and stay in their own sections.
3605    #[test]
3606    fn test_dns_outgoing_all_sections_spill() {
3607        let mut out = DnsOutgoing::new(FLAGS_QR_RESPONSE);
3608        for i in 0..40 {
3609            out.add_answer_record(ptr_answer(i));
3610        }
3611        for i in 40..80 {
3612            out.add_authority(Box::new(ptr_answer(i)));
3613        }
3614        for i in 80..120 {
3615            out.add_additional_answer(ptr_answer(i));
3616        }
3617
3618        let packets = out.to_packets(MAX_PKT_DEFAULT, IPV6);
3619        assert!(packets.len() > 1);
3620
3621        let mut answers = 0;
3622        let mut authorities = 0;
3623        let mut additionals = 0;
3624        for packet in &packets {
3625            assert!(packet.size() <= MAX_PKT_DEFAULT);
3626            let parsed = DnsIncoming::new(packet.as_bytes().to_vec(), test_interface_id()).unwrap();
3627            answers += parsed.answers().len();
3628            authorities += parsed.authorities().len();
3629            additionals += parsed.additionals().len();
3630        }
3631
3632        assert_eq!(answers, 40);
3633        assert_eq!(authorities, 40);
3634        assert_eq!(additionals, 40);
3635    }
3636    #[test]
3637    fn test_nsec_encoding_round_trip() {
3638        use super::DnsRecordExt;
3639
3640        for (bitmap, types) in [
3641            (vec![0x40], vec![1]),
3642            (vec![0, 0, 0, 8], vec![28]),
3643            (vec![0x40, 0, 0, 8], vec![1, 28]),
3644        ] {
3645            let mut out = DnsOutgoing::new(FLAGS_QR_RESPONSE | super::FLAGS_AA);
3646            out.add_answer_record(super::DnsNSec::new(
3647                "negative.local.",
3648                CLASS_IN | super::CLASS_CACHE_FLUSH,
3649                120,
3650                "negative.local.".to_string(),
3651                bitmap.clone(),
3652            ));
3653            let packets = out.to_packets(MAX_PKT_DEFAULT, IPV6);
3654            assert_eq!(packets.len(), 1);
3655            let data = packets[0].as_bytes().to_vec();
3656            // Owner name, RR header, compressed Next Domain Name, then window 0.
3657            let rdata_offset = 12 + 16 + 10;
3658            assert_eq!(
3659                &data[rdata_offset..rdata_offset + 4],
3660                &[0xc0, 0x0c, 0, bitmap.len() as u8]
3661            );
3662            assert_eq!(&data[rdata_offset + 4..], bitmap.as_slice());
3663            let incoming = DnsIncoming::new(data, test_interface_id()).unwrap();
3664            assert_eq!(incoming.answers().len(), 1);
3665            let record = incoming.answers()[0]
3666                .any()
3667                .downcast_ref::<super::DnsNSec>()
3668                .unwrap();
3669            assert_eq!(record.next_domain, "negative.local.");
3670            assert_eq!(record._types(), types);
3671            assert_eq!(record.get_record().get_ttl(), 120);
3672        }
3673    }
3674
3675    #[test]
3676    fn test_service_binding_questions_round_trip() {
3677        let mut out = DnsOutgoing::new(FLAGS_QR_QUERY);
3678        for ty in [RRType::SVCB, RRType::HTTPS, RRType::AAAA, RRType::A] {
3679            out.add_question("binding.local.", ty);
3680        }
3681        let packets = out.to_packets(MAX_PKT_DEFAULT, IPV6);
3682        let incoming =
3683            DnsIncoming::new(packets[0].as_bytes().to_vec(), test_interface_id()).unwrap();
3684        let types: Vec<_> = incoming.questions().iter().map(|q| q.entry.ty).collect();
3685        assert_eq!(
3686            types,
3687            [RRType::SVCB, RRType::HTTPS, RRType::AAAA, RRType::A]
3688        );
3689    }
3690}