Skip to main content

consomme/dns_resolver/
static_records.rs

1// Copyright (c) Microsoft Corporation.
2// Licensed under the MIT License.
3
4use smoltcp::wire::DnsFlags;
5use smoltcp::wire::DnsPacket;
6use smoltcp::wire::DnsQueryType;
7use smoltcp::wire::DnsQuestion;
8use thiserror::Error;
9
10/// DNS record type and data for a static record.
11///
12/// Only [`StaticDnsRecord::A`] is currently supported.
13#[derive(Debug, Clone, Copy, PartialEq, Eq)]
14pub enum StaticDnsRecord {
15    /// IPv4 host address.
16    A([u8; 4]),
17}
18
19/// An error adding a static DNS record.
20#[derive(Debug, Error, PartialEq, Eq)]
21pub enum StaticDnsRecordError {
22    /// The query name is empty, too long, or malformed.
23    #[error("the query name is empty, too long, or malformed")]
24    InvalidName,
25}
26
27/// DNS `CLASS` value for the Internet (`IN`) class.
28const DNS_CLASS_IN: u16 = 1;
29
30/// TTL advertised for static records.
31const DEFAULT_TTL: u32 = 60;
32
33/// Maximum length of a single DNS label, in bytes (RFC 1035 §2.3.4).
34const MAX_LABEL_LEN: usize = 63;
35
36/// Length of the fixed DNS message header, in bytes.
37const DNS_HEADER_LEN: usize = 12;
38
39/// Fixed per-answer overhead for a compression-pointer `A` record: name
40/// pointer (2) + TYPE (2) + CLASS (2) + TTL (4) + RDLENGTH (2).
41const ANSWER_FIXED_LEN: usize = 12;
42
43/// Maximum size of a DNS response over UDP.
44pub(crate) const MAX_DNS_UDP_RESPONSE_LEN: usize = 512;
45
46/// A single static DNS record.
47struct StaticDnsRecordEntry {
48    /// Lowercased presentation-form domain name (no trailing dot).
49    name: String,
50    record: StaticDnsRecord,
51}
52
53#[derive(Default)]
54pub struct StaticDnsRecords {
55    records: Vec<StaticDnsRecordEntry>,
56}
57
58impl StaticDnsRecords {
59    pub(super) fn is_empty(&self) -> bool {
60        self.records.is_empty()
61    }
62
63    /// Adds a static record.
64    ///
65    /// `name` is the query name in presentation form (e.g. `"example.com"`),
66    /// stored lowercased and compared case-insensitively. It must be ASCII.
67    ///
68    /// Returns [`StaticDnsRecordError::InvalidName`] if `name` is empty, too
69    /// long, non-ASCII, or otherwise malformed.
70    pub fn add(&mut self, record: StaticDnsRecord, name: &str) -> Result<(), StaticDnsRecordError> {
71        let name = normalize_name(name).ok_or(StaticDnsRecordError::InvalidName)?;
72        self.records.push(StaticDnsRecordEntry { name, record });
73        Ok(())
74    }
75
76    /// Builds a DNS response for `query` if it matches one of the static
77    /// records, otherwise returns `None`.
78    ///
79    /// `max_len` bounds the size of the returned DNS message (in bytes).
80    pub fn build_response(&self, query: &[u8], max_len: usize) -> Option<Vec<u8>> {
81        if self.records.is_empty() {
82            return None;
83        }
84
85        let packet = DnsPacket::new_checked(query).ok()?;
86
87        // Ignore packets with the response flag set.
88        if packet.flags().contains(DnsFlags::RESPONSE) {
89            return None;
90        }
91
92        let mut rest = packet.payload();
93        let mut answers = Vec::new();
94        for _ in 0..packet.question_count() {
95            let question_offset = query.len() - rest.len();
96            // `Question::parse` also validates that the class is `IN`.
97            let (next, question) = DnsQuestion::parse(rest).ok()?;
98            let qname = decode_name(&packet, question.name)?;
99            if question.type_ == DnsQueryType::A {
100                answers.extend(self.records.iter().filter_map(|rec| match &rec.record {
101                    StaticDnsRecord::A(address) if rec.name == qname => Some(StaticAnswer {
102                        name: if question_offset <= 0x3fff {
103                            AnswerName::Pointer(question_offset as u16)
104                        } else {
105                            AnswerName::Wire(question.name)
106                        },
107                        rdata: address.as_slice(),
108                    }),
109                    StaticDnsRecord::A(_) => None,
110                }));
111            }
112            rest = next;
113        }
114
115        if answers.is_empty() {
116            return None;
117        }
118
119        let question_section_len = packet.payload().len() - rest.len();
120        let question_section = &packet.payload()[..question_section_len];
121        let mut total = DNS_HEADER_LEN + question_section.len();
122
123        if total > max_len {
124            tracelimit::warn_ratelimited!(
125                required_len = total,
126                max_len,
127                "static DNS response buffer is too small for the header and question"
128            );
129            return None;
130        }
131
132        let mut fit = 0;
133        for answer in &answers {
134            let answer_len = answer.buffer_len();
135            if fit == u16::MAX as usize || total + answer_len > max_len {
136                break;
137            }
138
139            total += answer_len;
140            fit += 1;
141        }
142
143        let truncated = fit < answers.len();
144        Some(build_a_response(
145            &packet,
146            question_section,
147            &answers[..fit],
148            truncated,
149        ))
150    }
151}
152
153fn normalize_name(name: &str) -> Option<String> {
154    let name = name.strip_suffix('.').unwrap_or(name);
155    if name.is_empty() || name.len() > smoltcp::config::DNS_MAX_NAME_SIZE {
156        return None;
157    }
158
159    // Reject non-ASCII names.
160    if !name.is_ascii() {
161        return None;
162    }
163
164    // Reject empty labels ("..") and labels longer than the DNS maximum (63).
165    if name
166        .split('.')
167        .any(|label| label.is_empty() || label.len() > MAX_LABEL_LEN)
168    {
169        return None;
170    }
171
172    Some(name.to_ascii_lowercase())
173}
174
175/// Decodes a DNS name into lowercased presentation form (no trailing dot).
176///
177/// Returns `None` on malformed input, or if the name exceeds
178/// [`smoltcp::config::DNS_MAX_NAME_SIZE`].
179fn decode_name(packet: &DnsPacket<&[u8]>, name: &[u8]) -> Option<String> {
180    let mut qname = String::new();
181    for label in packet.parse_name(name) {
182        let label = label.ok()?;
183        if !label.is_ascii() || label.contains(&b'.') {
184            return None;
185        }
186        if !qname.is_empty() {
187            qname.push('.');
188        }
189        for &b in label {
190            qname.push(b.to_ascii_lowercase() as char);
191        }
192        if qname.len() > smoltcp::config::DNS_MAX_NAME_SIZE {
193            return None;
194        }
195    }
196    Some(qname)
197}
198
199enum AnswerName<'a> {
200    Pointer(u16),
201    Wire(&'a [u8]),
202}
203
204struct StaticAnswer<'a> {
205    name: AnswerName<'a>,
206    rdata: &'a [u8],
207}
208
209impl StaticAnswer<'_> {
210    fn buffer_len(&self) -> usize {
211        let name_len = match self.name {
212            AnswerName::Pointer(_) => 2,
213            AnswerName::Wire(name) => name.len(),
214        };
215        name_len + ANSWER_FIXED_LEN - 2 + self.rdata.len()
216    }
217}
218
219/// Builds a DNS response message containing one `A` answer per entry in
220/// `answers`, echoing the query's `question` section after the header.
221fn build_a_response(
222    query: &DnsPacket<&[u8]>,
223    question_section: &[u8],
224    answers: &[StaticAnswer<'_>],
225    truncated: bool,
226) -> Vec<u8> {
227    let ancount = answers.len().min(u16::MAX as usize) as u16;
228
229    // Response flags: QR=1, AA=1, RA=1, RD echoed from the query. TC is set
230    // when answers were dropped to fit the response-size budget.
231    let mut flags = DnsFlags::RESPONSE | DnsFlags::AUTHORITATIVE | DnsFlags::RECURSION_AVAILABLE;
232    flags |= query.flags() & DnsFlags::RECURSION_DESIRED;
233    if truncated {
234        flags |= DnsFlags::TRUNCATED;
235    }
236
237    // Header + echoed question section, written via smoltcp.
238    let mut response = vec![0u8; DNS_HEADER_LEN + question_section.len()];
239    {
240        let mut packet = DnsPacket::new_unchecked(&mut response[..]);
241        packet.set_transaction_id(query.transaction_id());
242        packet.set_flags(flags);
243        packet.set_opcode(query.opcode());
244        packet.set_question_count(query.question_count());
245        packet.set_answer_record_count(ancount);
246        packet.set_authority_record_count(0);
247        packet.set_additional_record_count(0);
248        packet.payload_mut().copy_from_slice(question_section);
249    }
250
251    for answer in answers.iter().take(ancount as usize) {
252        match answer.name {
253            AnswerName::Pointer(offset) => {
254                response.extend_from_slice(&(0xc000 | offset).to_be_bytes());
255            }
256            AnswerName::Wire(name) => response.extend_from_slice(name),
257        }
258        response.extend_from_slice(&u16::from(DnsQueryType::A).to_be_bytes());
259        response.extend_from_slice(&DNS_CLASS_IN.to_be_bytes());
260        response.extend_from_slice(&DEFAULT_TTL.to_be_bytes());
261        response.extend_from_slice(&(answer.rdata.len() as u16).to_be_bytes());
262        response.extend_from_slice(answer.rdata);
263    }
264
265    response
266}
267
268/// Builds a DNS query for `name` with the given qtype, in wire format.
269///
270/// Uses smoltcp's [`DnsRepr`] emitter, which always encodes the `IN` class.
271/// Shared by the static-record and DNS-over-TCP unit tests.
272#[cfg(test)]
273pub(crate) fn build_query(id: u16, name: &str, qtype: DnsQueryType) -> Vec<u8> {
274    use smoltcp::wire::DnsOpcode;
275    use smoltcp::wire::DnsRepr;
276
277    // Encode the query name into DNS wire format (length-prefixed labels).
278    let mut name_wire = Vec::new();
279    for label in name.split('.').filter(|l| !l.is_empty()) {
280        name_wire.push(label.len() as u8);
281        name_wire.extend_from_slice(label.as_bytes());
282    }
283
284    name_wire.push(0);
285
286    let repr = DnsRepr {
287        transaction_id: id,
288        opcode: DnsOpcode::Query,
289        flags: DnsFlags::RECURSION_DESIRED,
290        question: DnsQuestion {
291            name: &name_wire,
292            type_: qtype,
293        },
294    };
295    let mut buffer = vec![0u8; repr.buffer_len()];
296    repr.emit(&mut DnsPacket::new_unchecked(&mut buffer[..]));
297    buffer
298}
299
300#[cfg(test)]
301mod tests {
302    use super::*;
303
304    #[test]
305    fn add_and_match_a_record() {
306        let mut records = StaticDnsRecords::default();
307        records
308            .add(StaticDnsRecord::A([10, 0, 0, 5]), "Example.com")
309            .unwrap();
310
311        let query = build_query(0x1234, "example.com", DnsQueryType::A);
312        let response = records
313            .build_response(&query, MAX_DNS_UDP_RESPONSE_LEN)
314            .expect("should match");
315
316        // Transaction ID preserved.
317        assert_eq!(&response[0..2], &[0x12, 0x34]);
318
319        // QR + AA set, RD preserved, RA set, RCODE 0.
320        assert_eq!(response[2], 0x85);
321        assert_eq!(response[3], 0x80);
322
323        // ANCOUNT == 1.
324        assert_eq!(u16::from_be_bytes([response[6], response[7]]), 1);
325
326        // Final 4 RDATA bytes are the address we registered.
327        assert_eq!(&response[response.len() - 4..], &[10, 0, 0, 5]);
328
329        // RDATA is preceded by RDLENGTH == 4.
330        assert_eq!(
331            u16::from_be_bytes([response[response.len() - 6], response[response.len() - 5]]),
332            4
333        );
334    }
335
336    #[test]
337    fn case_insensitive_match() {
338        let mut records = StaticDnsRecords::default();
339        records
340            .add(StaticDnsRecord::A([1, 2, 3, 4]), "host.local")
341            .unwrap();
342        let query = build_query(1, "HOST.LOCAL", DnsQueryType::A);
343        assert!(
344            records
345                .build_response(&query, MAX_DNS_UDP_RESPONSE_LEN)
346                .is_some()
347        );
348    }
349
350    #[test]
351    fn multiple_records_same_name() {
352        let mut records = StaticDnsRecords::default();
353        records
354            .add(StaticDnsRecord::A([1, 1, 1, 1]), "many.test")
355            .unwrap();
356        records
357            .add(StaticDnsRecord::A([2, 2, 2, 2]), "many.test")
358            .unwrap();
359        let query = build_query(1, "many.test", DnsQueryType::A);
360        let response = records
361            .build_response(&query, MAX_DNS_UDP_RESPONSE_LEN)
362            .unwrap();
363        assert_eq!(u16::from_be_bytes([response[6], response[7]]), 2);
364    }
365
366    #[test]
367    fn any_matching_question_is_answered() {
368        let mut records = StaticDnsRecords::default();
369        records
370            .add(StaticDnsRecord::A([1, 2, 3, 4]), "known.test")
371            .unwrap();
372
373        let mut query = build_query(1, "unknown.test", DnsQueryType::A);
374        let matching_query = build_query(1, "known.test", DnsQueryType::A);
375        query[4..6].copy_from_slice(&2u16.to_be_bytes());
376        query.extend_from_slice(&matching_query[DNS_HEADER_LEN..]);
377
378        let response = records
379            .build_response(&query, MAX_DNS_UDP_RESPONSE_LEN)
380            .expect("one matching question should produce a response");
381
382        assert_eq!(u16::from_be_bytes([response[4], response[5]]), 2);
383        assert_eq!(u16::from_be_bytes([response[6], response[7]]), 1);
384        assert_eq!(&response[response.len() - 4..], &[1, 2, 3, 4]);
385    }
386
387    #[test]
388    fn non_matching_name_returns_none() {
389        let mut records = StaticDnsRecords::default();
390        records
391            .add(StaticDnsRecord::A([1, 2, 3, 4]), "known.test")
392            .unwrap();
393        let query = build_query(1, "unknown.test", DnsQueryType::A);
394        assert!(
395            records
396                .build_response(&query, MAX_DNS_UDP_RESPONSE_LEN)
397                .is_none()
398        );
399    }
400
401    #[test]
402    fn non_a_query_returns_none() {
403        let mut records = StaticDnsRecords::default();
404        records
405            .add(StaticDnsRecord::A([1, 2, 3, 4]), "known.test")
406            .unwrap();
407        // AAAA for the same name should not be answered.
408        let query = build_query(1, "known.test", DnsQueryType::Aaaa);
409        assert!(
410            records
411                .build_response(&query, MAX_DNS_UDP_RESPONSE_LEN)
412                .is_none()
413        );
414    }
415
416    #[test]
417    fn empty_store_returns_none() {
418        let records = StaticDnsRecords::default();
419        let query = build_query(1, "known.test", DnsQueryType::A);
420        assert!(
421            records
422                .build_response(&query, MAX_DNS_UDP_RESPONSE_LEN)
423                .is_none()
424        );
425    }
426
427    #[test]
428    fn malformed_queries_do_not_panic() {
429        let mut records = StaticDnsRecords::default();
430        records
431            .add(StaticDnsRecord::A([1, 2, 3, 4]), "known.test")
432            .unwrap();
433
434        // Too short, truncated label, unterminated name, compression pointer.
435        assert!(
436            records
437                .build_response(&[], MAX_DNS_UDP_RESPONSE_LEN)
438                .is_none()
439        );
440        assert!(
441            records
442                .build_response(&[0; 5], MAX_DNS_UDP_RESPONSE_LEN)
443                .is_none()
444        );
445
446        let mut truncated = build_query(1, "known.test", DnsQueryType::A);
447        truncated.truncate(15);
448
449        assert!(
450            records
451                .build_response(&truncated, MAX_DNS_UDP_RESPONSE_LEN)
452                .is_none()
453        );
454
455        // A label length that runs off the end of the buffer.
456        let bad = [0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 63, b'x'];
457        assert!(
458            records
459                .build_response(&bad, MAX_DNS_UDP_RESPONSE_LEN)
460                .is_none()
461        );
462
463        // Compression pointer in the question.
464        let ptr = [0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0xc0, 0x0c];
465        assert!(
466            records
467                .build_response(&ptr, MAX_DNS_UDP_RESPONSE_LEN)
468                .is_none()
469        );
470    }
471
472    #[test]
473    fn unrepresentable_wire_names_do_not_match() {
474        let mut records = StaticDnsRecords::default();
475        records
476            .add(StaticDnsRecord::A([1, 2, 3, 4]), "a.b")
477            .unwrap();
478
479        // A single wire-format label containing a dot must not be confused
480        // with two presentation-form labels.
481        let mut dotted_label = build_query(1, "axb", DnsQueryType::A);
482        dotted_label[14] = b'.';
483        assert!(
484            records
485                .build_response(&dotted_label, MAX_DNS_UDP_RESPONSE_LEN)
486                .is_none()
487        );
488
489        // Non-ASCII wire-format labels cannot represent names accepted by
490        // StaticDnsRecords::add.
491        let mut non_ascii_label = build_query(1, "a", DnsQueryType::A);
492        non_ascii_label[13] = 0xff;
493        let packet = DnsPacket::new_checked(non_ascii_label.as_slice()).unwrap();
494        let (_, question) = DnsQuestion::parse(packet.payload()).unwrap();
495        assert_eq!(decode_name(&packet, question.name), None);
496    }
497
498    #[test]
499    fn response_packets_are_ignored() {
500        let mut records = StaticDnsRecords::default();
501        records
502            .add(StaticDnsRecord::A([1, 2, 3, 4]), "known.test")
503            .unwrap();
504
505        // A matching query is answered...
506        let mut query = build_query(1, "known.test", DnsQueryType::A);
507        assert!(
508            records
509                .build_response(&query, MAX_DNS_UDP_RESPONSE_LEN)
510                .is_some()
511        );
512
513        // ...but the same message with the QR (response) bit set is ignored,
514        // so we don't reply to a DNS response misrouted to port 53.
515        query[2] |= 0x80;
516        assert!(
517            records
518                .build_response(&query, MAX_DNS_UDP_RESPONSE_LEN)
519                .is_none()
520        );
521    }
522
523    #[test]
524    fn oversized_answer_set_is_truncated() {
525        let mut records = StaticDnsRecords::default();
526        // Register more answers than a small budget can hold.
527        for i in 0..20u8 {
528            records
529                .add(StaticDnsRecord::A([10, 0, 0, i]), "many.test")
530                .unwrap();
531        }
532
533        let query = build_query(1, "many.test", DnsQueryType::A);
534
535        // With a generous budget all answers fit and TC is clear.
536        let full = records
537            .build_response(&query, MAX_DNS_UDP_RESPONSE_LEN)
538            .unwrap();
539        let full_count = u16::from_be_bytes([full[6], full[7]]);
540        assert_eq!(full_count, 20);
541        assert_eq!(full[2] & 0x02, 0, "TC must be clear when nothing dropped");
542
543        // Size the budget to hold exactly two answers.
544        let answer_len = ANSWER_FIXED_LEN + 4;
545        let base = full.len() - full_count as usize * answer_len;
546        let budget = base + answer_len * 2;
547        let response = records.build_response(&query, budget).unwrap();
548
549        assert_eq!(u16::from_be_bytes([response[6], response[7]]), 2);
550        assert_ne!(response[2] & 0x02, 0, "TC must be set when answers dropped");
551        assert!(response.len() <= budget);
552    }
553
554    #[test]
555    fn answer_count_is_truncated_to_header_limit() {
556        let mut records = StaticDnsRecords::default();
557        for _ in 0..=u16::MAX {
558            records
559                .add(StaticDnsRecord::A([10, 0, 0, 1]), "many.test")
560                .unwrap();
561        }
562
563        let query = build_query(1, "many.test", DnsQueryType::A);
564        let response = records.build_response(&query, usize::MAX).unwrap();
565
566        assert_eq!(u16::from_be_bytes([response[6], response[7]]), u16::MAX);
567        assert_ne!(response[2] & 0x02, 0, "TC must be set when answers dropped");
568    }
569
570    #[test]
571    fn budget_too_small_for_header_returns_none() {
572        let mut records = StaticDnsRecords::default();
573        records
574            .add(StaticDnsRecord::A([1, 2, 3, 4]), "known.test")
575            .unwrap();
576        let query = build_query(1, "known.test", DnsQueryType::A);
577        // Not even the header and question fit, so no response is synthesized.
578        assert!(records.build_response(&query, 4).is_none());
579    }
580
581    #[test]
582    fn add_validation() {
583        let mut records = StaticDnsRecords::default();
584
585        // Empty name.
586        assert_eq!(
587            records.add(StaticDnsRecord::A([1, 2, 3, 4]), ""),
588            Err(StaticDnsRecordError::InvalidName)
589        );
590    }
591
592    #[test]
593    fn add_rejects_malformed_names() {
594        let mut records = StaticDnsRecords::default();
595
596        // Consecutive dots ("..") produce an empty label.
597        assert_eq!(
598            records.add(StaticDnsRecord::A([1, 2, 3, 4]), "a..b"),
599            Err(StaticDnsRecordError::InvalidName)
600        );
601
602        // A leading dot is also an empty label.
603        assert_eq!(
604            records.add(StaticDnsRecord::A([1, 2, 3, 4]), ".example.com"),
605            Err(StaticDnsRecordError::InvalidName)
606        );
607
608        // A name longer than the maximum permitted length is rejected.
609        let too_long = "a".repeat(smoltcp::config::DNS_MAX_NAME_SIZE + 1);
610        assert_eq!(
611            records.add(StaticDnsRecord::A([1, 2, 3, 4]), &too_long),
612            Err(StaticDnsRecordError::InvalidName)
613        );
614
615        // A single label longer than the DNS maximum (63) is rejected, even
616        // when the overall name is within the length limit.
617        let long_label = "a".repeat(MAX_LABEL_LEN + 1);
618        assert_eq!(
619            records.add(StaticDnsRecord::A([1, 2, 3, 4]), &long_label),
620            Err(StaticDnsRecordError::InvalidName)
621        );
622
623        // A non-ASCII name is rejected.
624        assert_eq!(
625            records.add(StaticDnsRecord::A([1, 2, 3, 4]), "exämple.com"),
626            Err(StaticDnsRecordError::InvalidName)
627        );
628
629        // A label of exactly the maximum length is accepted.
630        let max_label = "a".repeat(MAX_LABEL_LEN);
631        assert!(
632            records
633                .add(StaticDnsRecord::A([1, 2, 3, 4]), &max_label)
634                .is_ok()
635        );
636
637        // A well-formed name (with an optional trailing dot) succeeds.
638        assert!(
639            records
640                .add(StaticDnsRecord::A([1, 2, 3, 4]), "valid.example.com.")
641                .is_ok()
642        );
643    }
644}