1use std::{collections::HashMap, fmt};
2
3use bytes::BufMut;
4
5use super::{
6 header::{Header, be_header},
7 name::{NameCompression, put_name},
8 question::{QueryClass, QueryType, Question, be_question},
9 record::{
10 Class, RData, ResourceRecord, Type, be_record,
11 endpoint::{EndpointAddr, WriteEndpointAddr},
12 srv::Srv,
13 },
14};
15use crate::core::parser::header::WriteHeader;
16
17#[derive(Default, Clone)]
19pub struct Packet {
20 pub header: Header,
21 pub questions: Vec<Question>,
22 pub answers: Vec<ResourceRecord>,
23 pub nameservers: Vec<ResourceRecord>,
24 pub additional: Vec<ResourceRecord>,
25}
26
27impl fmt::Display for Packet {
28 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
29 writeln!(f, "DNS Packet:")?;
30 writeln!(
31 f,
32 " Header: ID={}, QR={}, AA={}, RCODE={:?}",
33 self.header.id,
34 self.header.flags.query(),
35 self.header.flags.authoritative(),
36 self.header.flags.response_code()
37 )?;
38 if !self.questions.is_empty() {
39 writeln!(f, " Questions:")?;
40 for q in &self.questions {
41 writeln!(f, " {} {:?} {:?}", q.name, q.qclass, q.qtype)?;
42 }
43 }
44 if !self.answers.is_empty() {
45 writeln!(f, " Answers:")?;
46 for rr in &self.answers {
47 write!(f, " {} {} {:?} {:?}", rr.name, rr.ttl, rr.cls, rr.typ)?;
48 match &rr.data {
49 RData::A(ip) => writeln!(f, " A {}", ip)?,
50 RData::AAAA(ip) => writeln!(f, " AAAA {}", ip)?,
51 RData::CName(name) => writeln!(f, " CNAME {}", name)?,
52 RData::E(ep) => {
53 writeln!(f, " E {}", ep)?;
54 }
55 _ => writeln!(f, " {:?}", rr.data)?,
56 }
57 }
58 }
59 if !self.nameservers.is_empty() {
60 writeln!(f, " Nameservers:")?;
61 for rr in &self.nameservers {
62 writeln!(f, " {} {} {:?} {:?}", rr.name, rr.ttl, rr.cls, rr.typ)?;
63 }
64 }
65 if !self.additional.is_empty() {
66 writeln!(f, " Additional:")?;
67 for rr in &self.additional {
68 writeln!(f, " {} {} {:?} {:?}", rr.name, rr.ttl, rr.cls, rr.typ)?;
69 }
70 }
71 Ok(())
72 }
73}
74
75impl Packet {
76 pub fn id(&self) -> u16 {
77 self.header.id
78 }
79
80 pub fn is_query(&self) -> bool {
81 self.header.flags.query()
82 }
83
84 pub fn query_with_id(service_name: String) -> Self {
85 let mut packet = Packet::default();
86 let id: u16 = rand::random();
87 packet.header.id = id;
88 packet.header.flags.set_query(false);
89 packet.add_question(&service_name, QueryType::A, QueryClass::IN, false);
90 packet
91 }
92
93 pub fn query(service_name: String) -> Self {
94 let mut packet = Self::default();
95 packet.add_question(&service_name, QueryType::A, QueryClass::IN, true);
96 packet
97 }
98
99 pub fn answer(id: u16, hosts: &HashMap<String, Vec<EndpointAddr>>) -> Self {
100 let mut packet = Self::default();
101 packet.header.id = id;
102 packet.header.flags.set_query(true);
103 hosts.iter().for_each(|(name, eps)| {
104 eps.iter().for_each(|ep| {
105 let (rtype, rdata) = (Type::E, RData::E(ep.clone()));
106 packet.add_answer(name, rtype, Class::IN, 300, rdata);
107 });
108 });
109 packet
110 }
111
112 pub fn to_bytes(&self) -> Vec<u8> {
113 let mut buf = Vec::with_capacity(2048);
114 let mut ctx = NameCompression::new();
115
116 buf.put_header(&self.header);
117
118 for question in &self.questions {
119 let _ = put_name(&mut buf, &question.name, &mut ctx);
120 buf.put_u16(question.qtype.into());
121 let mut qclass = u16::from(question.qclass);
122 if question.prefer_unicast {
123 qclass |= 0x8000;
124 }
125 buf.put_u16(qclass);
126 }
127
128 for answer in &self.answers {
129 put_record(&mut buf, answer, &mut ctx);
130 }
131 for nameserver in &self.nameservers {
132 put_record(&mut buf, nameserver, &mut ctx);
133 }
134 for additional in &self.additional {
135 put_record(&mut buf, additional, &mut ctx);
136 }
137
138 buf
139 }
140
141 fn add_question(
142 &mut self,
143 qname: &str,
144 qtype: QueryType,
145 qclass: QueryClass,
146 prefer_unicast: bool,
147 ) {
148 let question = Question {
149 name: qname.to_string(),
150 prefer_unicast,
151 qtype,
152 qclass,
153 };
154 self.header.questions_count += 1;
155 self.questions.push(question);
156 }
157
158 fn add_answer(&mut self, name: &str, rtype: Type, rclass: Class, ttl: u32, data: RData) {
159 let response = ResourceRecord {
160 name: name.to_string(),
161 typ: rtype,
162 multicast_unique: false,
163 cls: rclass,
164 ttl,
165 data,
166 };
167 self.header.flags.set_query(true);
169 self.header.answers_count += 1;
170 self.answers.push(response);
171 }
172}
173
174fn put_record(buf: &mut Vec<u8>, record: &ResourceRecord, ctx: &mut NameCompression) {
175 let _ = put_name(buf, &record.name, ctx);
176 buf.put_u16(u16::from(record.typ));
177
178 let mut cls = u16::from(record.cls);
179 if record.multicast_unique {
180 cls |= 0x8000;
181 }
182 buf.put_u16(cls);
183
184 buf.put_u32(record.ttl);
185
186 let rdlen_pos = buf.len();
187 buf.put_u16(0);
188 let rdata_start = buf.len();
189
190 match &record.data {
191 RData::A(ip) => buf.put_slice(&ip.octets()),
192 RData::AAAA(ip) => buf.put_slice(&ip.octets()),
193 RData::CName(name) => {
194 let _ = put_name(buf, name, ctx);
195 }
196 RData::Txt(txt) => buf.put_slice(txt),
197 RData::Srv(srv) => put_srv(buf, srv, ctx),
198 RData::Ptr(ptr) => {
199 let _ = put_name(buf, ptr.name(), ctx);
200 }
201 RData::E(e) => buf.put_endpoint_addr(e),
202 }
203
204 let rdlen = (buf.len() - rdata_start) as u16;
205 let [hi, lo] = rdlen.to_be_bytes();
206 buf[rdlen_pos] = hi;
207 buf[rdlen_pos + 1] = lo;
208}
209
210fn put_srv(buf: &mut Vec<u8>, srv: &Srv, ctx: &mut NameCompression) {
211 buf.put_u16(srv.priority());
212 buf.put_u16(srv.weight());
213 buf.put_u16(srv.port());
214 let _ = put_name(buf, srv.target(), ctx);
215}
216
217pub fn be_packet(input: &[u8]) -> nom::IResult<&[u8], Packet> {
218 let (remain, header) = be_header(input)?;
219 let (remain, ()) = validate_section_counts(remain, &header)?;
220
221 let (remain, questions) =
222 parse::<Question>(remain, input, header.questions_count, be_question)?;
223 let (remain, answers) =
224 parse::<ResourceRecord>(remain, input, header.answers_count, be_record)?;
225 let (remain, nameservers) =
226 parse::<ResourceRecord>(remain, input, header.nameservers_count, be_record)?;
227 let (remain, additional) =
228 parse::<ResourceRecord>(remain, input, header.additional_count, be_record)?;
229
230 Ok((
231 remain,
232 Packet {
233 header,
234 questions,
235 answers,
236 nameservers,
237 additional,
238 },
239 ))
240}
241
242const MIN_QUESTION_WIRE_LEN: usize = 5;
244const MIN_RESOURCE_RECORD_WIRE_LEN: usize = 11;
246
247fn validate_section_counts<'a>(input: &'a [u8], header: &Header) -> nom::IResult<&'a [u8], ()> {
249 let question_bytes = usize::from(header.questions_count).checked_mul(MIN_QUESTION_WIRE_LEN);
250 let record_count = usize::from(header.answers_count)
251 .checked_add(usize::from(header.nameservers_count))
252 .and_then(|count| count.checked_add(usize::from(header.additional_count)));
253 let minimum_wire_len = question_bytes.and_then(|question_bytes| {
254 record_count
255 .and_then(|count| count.checked_mul(MIN_RESOURCE_RECORD_WIRE_LEN))
256 .and_then(|record_bytes| question_bytes.checked_add(record_bytes))
257 });
258
259 let Some(minimum_wire_len) = minimum_wire_len else {
260 return Err(nom::Err::Failure(nom::error::Error::new(
261 input,
262 nom::error::ErrorKind::Verify,
263 )));
264 };
265
266 if minimum_wire_len > input.len() {
267 return Err(nom::Err::Failure(nom::error::Error::new(
268 input,
269 nom::error::ErrorKind::Verify,
270 )));
271 }
272
273 Ok((input, ()))
274}
275
276fn parse<'a, T>(
277 mut input: &'a [u8],
278 original: &'a [u8],
279 count: u16,
280 parser: impl Fn(&'a [u8], &'a [u8]) -> nom::IResult<&'a [u8], T>,
281) -> nom::IResult<&'a [u8], Vec<T>> {
282 let mut records = Vec::new();
283 for _ in 0..count {
284 match parser(input, original) {
285 Ok((new_input, record)) => {
286 records.push(record);
287 input = new_input;
288 }
289 Err(nom::Err::Error(nom::error::Error {
290 input: remaining, ..
291 })) => {
292 input = remaining;
293 }
294 _ => break,
295 }
296 }
297 Ok((input, records))
298}
299
300impl fmt::Debug for Packet {
301 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
302 write!(
303 f,
304 "Packet {{ id: {}, qr: {}, opcode: {:?}, rcode: {:?}, questions: {}, answers: {}, authorities: {}, additional: {} }}",
305 self.header.id,
306 self.header.flags.query(),
307 self.header.flags.opcode(),
308 self.header.flags.response_code(),
309 self.questions.len(),
310 self.answers.len(),
311 self.nameservers.len(),
312 self.additional.len()
313 )
314 }
315}
316
317#[cfg(test)]
318mod test {
319 use std::net::{IpAddr, Ipv4Addr, Ipv6Addr};
320
321 use super::*;
322 use crate::core::parser::{
323 self,
324 question::{QueryClass, QueryType},
325 record::{Class, RData, Type, srv::Srv},
326 };
327
328 fn decode_hex(s: &str) -> Vec<u8> {
329 let mut out = Vec::with_capacity(s.len() / 2);
330 let mut n = 0u8;
331 let mut high = true;
332 for b in s.bytes() {
333 let v = match b {
334 b'0'..=b'9' => b - b'0',
335 b'a'..=b'f' => b - b'a' + 10,
336 b'A'..=b'F' => b - b'A' + 10,
337 b' ' | b'\n' | b'\r' | b'\t' => continue,
338 _ => panic!("invalid hex"),
339 };
340 if high {
341 n = v << 4;
342 high = false;
343 } else {
344 n |= v;
345 out.push(n);
346 high = true;
347 }
348 }
349 if !high {
350 panic!("odd hex length");
351 }
352 out
353 }
354
355 #[test]
356 fn parse_example_query() {
357 let query = b"\x06%\x01\x00\x00\x01\x00\x00\x00\x00\x00\x00\
358 \x07example\x03com\x00\x00\x01\x00\x01";
359 let (_, packet) = be_packet(query).unwrap();
360 assert_eq!(packet.header.id, 1573);
361 assert_eq!(packet.header.questions_count, 1);
362 assert_eq!(packet.questions[0].qtype, QueryType::A);
363 assert_eq!(packet.questions[0].qclass, QueryClass::IN);
364 assert_eq!(packet.questions[0].name.to_string(), "example.com");
365 assert_eq!(packet.header.answers_count, 0);
366 }
367
368 #[test]
369 fn parse_example_response() {
370 let response = b"\x06%\x81\x80\x00\x01\x00\x01\x00\x00\x00\x00\
371 \x07example\x03com\x00\x00\x01\x00\x01\
372 \xc0\x0c\x00\x01\x00\x01\x00\x00\x04\xf8\
373 \x00\x04]\xb8\xd8\"";
374 let (_, packet) = be_packet(response).unwrap();
375 assert_eq!(packet.header.id, 1573);
376 assert_eq!(packet.header.questions_count, 1);
377 assert_eq!(packet.questions[0].qtype, QueryType::A);
378 assert_eq!(packet.questions[0].qclass, QueryClass::IN);
379 assert_eq!(packet.questions[0].name.to_string(), "example.com");
380 assert_eq!(packet.header.answers_count, 1);
381 assert_eq!(packet.answers[0].name.to_string(), "example.com");
382 assert_eq!(packet.answers[0].typ, Type::A);
383 assert_eq!(packet.answers[0].cls, Class::IN);
384 assert_eq!(packet.answers[0].ttl, 1272);
385 match &packet.answers[0].data {
386 RData::A(addr) => assert_eq!(*addr, Ipv4Addr::new(93, 184, 216, 34)),
387 _ => panic!("unexpected rdata"),
388 }
389 }
390
391 #[test]
392 fn parse_response_with_multicast_unique() {
393 let response = b"\x06%\x81\x80\x00\x01\x00\x01\x00\x00\x00\x00\
394 \x07example\x03com\x00\x00\x01\x00\x01\
395 \xc0\x0c\x00\x01\x80\x01\x00\x00\x04\xf8\
396 \x00\x04]\xb8\xd8\"";
397 let (_, packet) = be_packet(response).unwrap();
398
399 assert_eq!(packet.answers.len(), 1);
400 assert!(packet.answers[0].multicast_unique);
401 assert_eq!(packet.answers[0].cls, Class::IN);
402 }
403
404 #[test]
405 fn parse_additional_record_response() {
406 tracing_subscriber::fmt().with_ansi(false).init();
407 let response = b"\x4a\xf0\x81\x80\x00\x01\x00\x01\x00\x01\x00\x01\
408 \x03www\x05skype\x03com\x00\x00\x01\x00\x01\
409 \xc0\x0c\x00\x05\x00\x01\x00\x00\x0e\x10\
410 \x00\x1c\x07\x6c\x69\x76\x65\x63\x6d\x73\x0e\x74\
411 \x72\x61\x66\x66\x69\x63\x6d\x61\x6e\x61\x67\x65\
412 \x72\x03\x6e\x65\x74\x00\
413 \xc0\x42\x00\x02\x00\x01\x00\x01\xd5\xd3\x00\x11\
414 \x01\x67\x0c\x67\x74\x6c\x64\x2d\x73\x65\x72\x76\x65\x72\x73\
415 \xc0\x42\
416 \x01\x61\xc0\x55\x00\x01\x00\x01\x00\x00\xa3\x1c\
417 \x00\x04\xc0\x05\x06\x1e";
418 let (_, packet) = be_packet(response).unwrap();
419
420 assert_eq!(packet.header.id, 19184);
421 assert_eq!(packet.header.questions_count, 1);
422 assert_eq!(packet.questions[0].qtype, QueryType::A);
423 assert_eq!(packet.questions[0].qclass, QueryClass::IN);
424 assert_eq!(&packet.questions[0].name.to_string()[..], "www.skype.com");
425 assert_eq!(packet.answers.len(), 1);
426 assert_eq!(&packet.answers[0].name.to_string()[..], "www.skype.com");
427 assert_eq!(packet.answers[0].cls, Class::IN);
428 assert_eq!(packet.answers[0].ttl, 3600);
429
430 match &packet.answers[0].data {
431 RData::CName(cname) => {
432 assert_eq!(cname, "livecms.trafficmanager.net");
433 }
434 ref x => panic!("Wrong rdata {x:?}"),
435 }
436 assert_eq!(packet.additional.len(), 1);
437 assert_eq!(
438 &packet.additional[0].name.to_string()[..],
439 "a.gtld-servers.net"
440 );
441 assert_eq!(packet.additional[0].cls, Class::IN);
442 assert_eq!(packet.additional[0].ttl, 41756);
443 match packet.additional[0].data {
444 RData::A(addr) => {
445 assert_eq!(addr, Ipv4Addr::new(192, 5, 6, 30));
446 }
447 ref x => panic!("Wrong rdata {x:?}"),
448 }
449 }
450
451 #[test]
452 fn parse_pack_packet() {
453 let mut response = Packet::default();
454 let address = [
455 IpAddr::V4(Ipv4Addr::new(192, 168, 1, 1)),
456 IpAddr::V6(Ipv6Addr::new(
457 0x2001, 0x0db8, 0x85a3, 0x0000, 0x0000, 0x8a2e, 0x0370, 0x7334,
458 )),
459 ];
460 for ip in address.iter() {
461 let (rtype, rdata) = match ip {
462 IpAddr::V4(ipv4_addr) => (
463 parser::record::Type::A,
464 parser::record::RData::A(*ipv4_addr),
465 ),
466 IpAddr::V6(ipv6_addr) => (
467 parser::record::Type::AAAA,
468 parser::record::RData::AAAA(*ipv6_addr),
469 ),
470 };
471 response.add_answer("example.com", rtype, parser::record::Class::IN, 300, rdata);
472 }
473
474 let srv = Srv::new(0, 0, 6000, "example.com".to_string());
475 response.add_answer(
476 "example.com",
477 parser::record::Type::Srv,
478 parser::record::Class::IN,
479 300,
480 parser::record::RData::Srv(srv),
481 );
482 let packet = response.to_bytes();
483 let (_, parsed_packet) = be_packet(&packet).unwrap();
484 assert_eq!(parsed_packet.header.id, response.header.id);
485 assert_eq!(
486 parsed_packet.header.questions_count,
487 response.header.questions_count
488 );
489 assert_eq!(parsed_packet.answers.len(), response.answers.len());
490 assert_eq!(parsed_packet.nameservers.len(), response.nameservers.len());
491 assert_eq!(parsed_packet.additional.len(), response.additional.len());
492 assert_eq!(parsed_packet.answers[0].name, response.answers[0].name);
493 }
494
495 #[test]
496 fn malformed_packet_does_not_panic() {
497 let packet_hex = "0021641c0000000100000000000078787878787878787878787303636f6d0000100001";
498 let data = decode_hex(packet_hex);
499 let ret = std::panic::catch_unwind(|| {
500 let _ = be_packet(&data);
501 });
502 assert!(ret.is_ok());
503 }
504
505 #[test]
506 fn impossible_section_counts_are_rejected_before_parsing() {
507 let packet = [
508 0, 0, 0, 0, 0xff, 0xff, 0, 0, 0, 0, 0, 0, ];
515
516 assert!(matches!(be_packet(&packet), Err(nom::Err::Failure(_))));
517 }
518
519 #[test]
520 fn aggregate_resource_record_counts_use_the_input_bound() {
521 let packet = [
522 0, 0, 0, 0, 0, 0, 0, 1, 0, 1, 0, 1, ];
529
530 assert!(matches!(be_packet(&packet), Err(nom::Err::Failure(_))));
531 }
532
533 #[test]
534 fn packet_with_unknown_rr_type_does_not_panic() {
535 let packet_hex = "8116840000010001000000000569627a6c700474657374046d69656b026e6c00000a0001c00c000a0001000000000005497f000001";
536 let data = decode_hex(packet_hex);
537 let ret = std::panic::catch_unwind(|| be_packet(&data));
538 assert!(ret.is_ok());
539 let (_, packet) = be_packet(&data).unwrap();
540 assert_eq!(packet.header.questions_count, 1);
541 assert_eq!(packet.header.answers_count, 1);
542 assert!(packet.answers.is_empty());
543 }
544
545 #[test]
546 fn packet_to_bytes_uses_name_compression_across_sections() {
547 let mut packet = Packet::default();
548 packet.add_question("www.skype.com", QueryType::A, QueryClass::IN, false);
549 packet.add_answer(
550 "mail.skype.com",
551 Type::A,
552 Class::IN,
553 1,
554 RData::A(Ipv4Addr::new(1, 2, 3, 4)),
555 );
556
557 let bytes = packet.to_bytes();
558 let mail_pos = bytes
559 .windows(5)
560 .position(|w| w == b"\x04mail")
561 .expect("mail label missing");
562 let skype_pos = bytes
563 .windows(6)
564 .position(|w| w == b"\x05skype")
565 .expect("skype label missing");
566
567 let ptr = (0xC000u16 | (skype_pos as u16)).to_be_bytes();
568 assert_eq!(&bytes[mail_pos + 5..mail_pos + 7], &ptr);
569 }
570}