1use smoltcp::wire::DnsFlags;
5use smoltcp::wire::DnsPacket;
6use smoltcp::wire::DnsQueryType;
7use smoltcp::wire::DnsQuestion;
8use thiserror::Error;
9
10#[derive(Debug, Clone, Copy, PartialEq, Eq)]
14pub enum StaticDnsRecord {
15 A([u8; 4]),
17}
18
19#[derive(Debug, Error, PartialEq, Eq)]
21pub enum StaticDnsRecordError {
22 #[error("the query name is empty, too long, or malformed")]
24 InvalidName,
25}
26
27const DNS_CLASS_IN: u16 = 1;
29
30const DEFAULT_TTL: u32 = 60;
32
33const MAX_LABEL_LEN: usize = 63;
35
36const DNS_HEADER_LEN: usize = 12;
38
39const ANSWER_FIXED_LEN: usize = 12;
42
43pub(crate) const MAX_DNS_UDP_RESPONSE_LEN: usize = 512;
45
46struct StaticDnsRecordEntry {
48 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 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 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 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 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 if !name.is_ascii() {
161 return None;
162 }
163
164 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
175fn 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
219fn 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 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 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#[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 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 assert_eq!(&response[0..2], &[0x12, 0x34]);
318
319 assert_eq!(response[2], 0x85);
321 assert_eq!(response[3], 0x80);
322
323 assert_eq!(u16::from_be_bytes([response[6], response[7]]), 1);
325
326 assert_eq!(&response[response.len() - 4..], &[10, 0, 0, 5]);
328
329 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 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 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 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 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 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 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 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 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 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 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 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 assert!(records.build_response(&query, 4).is_none());
579 }
580
581 #[test]
582 fn add_validation() {
583 let mut records = StaticDnsRecords::default();
584
585 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 assert_eq!(
598 records.add(StaticDnsRecord::A([1, 2, 3, 4]), "a..b"),
599 Err(StaticDnsRecordError::InvalidName)
600 );
601
602 assert_eq!(
604 records.add(StaticDnsRecord::A([1, 2, 3, 4]), ".example.com"),
605 Err(StaticDnsRecordError::InvalidName)
606 );
607
608 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 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 assert_eq!(
625 records.add(StaticDnsRecord::A([1, 2, 3, 4]), "exämple.com"),
626 Err(StaticDnsRecordError::InvalidName)
627 );
628
629 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 assert!(
639 records
640 .add(StaticDnsRecord::A([1, 2, 3, 4]), "valid.example.com.")
641 .is_ok()
642 );
643 }
644}