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
220 let (remain, questions) =
221 parse::<Question>(remain, input, header.questions_count, be_question)?;
222 let (remain, answers) =
223 parse::<ResourceRecord>(remain, input, header.answers_count, be_record)?;
224 let (remain, nameservers) =
225 parse::<ResourceRecord>(remain, input, header.nameservers_count, be_record)?;
226 let (remain, additional) =
227 parse::<ResourceRecord>(remain, input, header.additional_count, be_record)?;
228
229 Ok((
230 remain,
231 Packet {
232 header,
233 questions,
234 answers,
235 nameservers,
236 additional,
237 },
238 ))
239}
240
241fn parse<'a, T>(
242 mut input: &'a [u8],
243 original: &'a [u8],
244 count: u16,
245 parser: impl Fn(&'a [u8], &'a [u8]) -> nom::IResult<&'a [u8], T>,
246) -> nom::IResult<&'a [u8], Vec<T>> {
247 let mut records = Vec::with_capacity(count as usize);
248 for _ in 0..count {
249 match parser(input, original) {
250 Ok((new_input, record)) => {
251 records.push(record);
252 input = new_input;
253 }
254 Err(nom::Err::Error(nom::error::Error {
255 input: remaining, ..
256 })) => {
257 input = remaining;
258 }
259 _ => break,
260 }
261 }
262 Ok((input, records))
263}
264
265impl fmt::Debug for Packet {
266 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
267 write!(
268 f,
269 "Packet {{ id: {}, qr: {}, opcode: {:?}, rcode: {:?}, questions: {}, answers: {}, authorities: {}, additional: {} }}",
270 self.header.id,
271 self.header.flags.query(),
272 self.header.flags.opcode(),
273 self.header.flags.response_code(),
274 self.questions.len(),
275 self.answers.len(),
276 self.nameservers.len(),
277 self.additional.len()
278 )
279 }
280}
281
282#[cfg(test)]
283mod test {
284 use std::net::{IpAddr, Ipv4Addr, Ipv6Addr};
285
286 use super::*;
287 use crate::core::parser::{
288 self,
289 question::{QueryClass, QueryType},
290 record::{Class, RData, Type, srv::Srv},
291 };
292
293 fn decode_hex(s: &str) -> Vec<u8> {
294 let mut out = Vec::with_capacity(s.len() / 2);
295 let mut n = 0u8;
296 let mut high = true;
297 for b in s.bytes() {
298 let v = match b {
299 b'0'..=b'9' => b - b'0',
300 b'a'..=b'f' => b - b'a' + 10,
301 b'A'..=b'F' => b - b'A' + 10,
302 b' ' | b'\n' | b'\r' | b'\t' => continue,
303 _ => panic!("invalid hex"),
304 };
305 if high {
306 n = v << 4;
307 high = false;
308 } else {
309 n |= v;
310 out.push(n);
311 high = true;
312 }
313 }
314 if !high {
315 panic!("odd hex length");
316 }
317 out
318 }
319
320 #[test]
321 fn parse_example_query() {
322 let query = b"\x06%\x01\x00\x00\x01\x00\x00\x00\x00\x00\x00\
323 \x07example\x03com\x00\x00\x01\x00\x01";
324 let (_, packet) = be_packet(query).unwrap();
325 assert_eq!(packet.header.id, 1573);
326 assert_eq!(packet.header.questions_count, 1);
327 assert_eq!(packet.questions[0].qtype, QueryType::A);
328 assert_eq!(packet.questions[0].qclass, QueryClass::IN);
329 assert_eq!(packet.questions[0].name.to_string(), "example.com");
330 assert_eq!(packet.header.answers_count, 0);
331 }
332
333 #[test]
334 fn parse_example_response() {
335 let response = b"\x06%\x81\x80\x00\x01\x00\x01\x00\x00\x00\x00\
336 \x07example\x03com\x00\x00\x01\x00\x01\
337 \xc0\x0c\x00\x01\x00\x01\x00\x00\x04\xf8\
338 \x00\x04]\xb8\xd8\"";
339 let (_, packet) = be_packet(response).unwrap();
340 assert_eq!(packet.header.id, 1573);
341 assert_eq!(packet.header.questions_count, 1);
342 assert_eq!(packet.questions[0].qtype, QueryType::A);
343 assert_eq!(packet.questions[0].qclass, QueryClass::IN);
344 assert_eq!(packet.questions[0].name.to_string(), "example.com");
345 assert_eq!(packet.header.answers_count, 1);
346 assert_eq!(packet.answers[0].name.to_string(), "example.com");
347 assert_eq!(packet.answers[0].typ, Type::A);
348 assert_eq!(packet.answers[0].cls, Class::IN);
349 assert_eq!(packet.answers[0].ttl, 1272);
350 match &packet.answers[0].data {
351 RData::A(addr) => assert_eq!(*addr, Ipv4Addr::new(93, 184, 216, 34)),
352 _ => panic!("unexpected rdata"),
353 }
354 }
355
356 #[test]
357 fn parse_response_with_multicast_unique() {
358 let response = b"\x06%\x81\x80\x00\x01\x00\x01\x00\x00\x00\x00\
359 \x07example\x03com\x00\x00\x01\x00\x01\
360 \xc0\x0c\x00\x01\x80\x01\x00\x00\x04\xf8\
361 \x00\x04]\xb8\xd8\"";
362 let (_, packet) = be_packet(response).unwrap();
363
364 assert_eq!(packet.answers.len(), 1);
365 assert!(packet.answers[0].multicast_unique);
366 assert_eq!(packet.answers[0].cls, Class::IN);
367 }
368
369 #[test]
370 fn parse_additional_record_response() {
371 tracing_subscriber::fmt().with_ansi(false).init();
372 let response = b"\x4a\xf0\x81\x80\x00\x01\x00\x01\x00\x01\x00\x01\
373 \x03www\x05skype\x03com\x00\x00\x01\x00\x01\
374 \xc0\x0c\x00\x05\x00\x01\x00\x00\x0e\x10\
375 \x00\x1c\x07\x6c\x69\x76\x65\x63\x6d\x73\x0e\x74\
376 \x72\x61\x66\x66\x69\x63\x6d\x61\x6e\x61\x67\x65\
377 \x72\x03\x6e\x65\x74\x00\
378 \xc0\x42\x00\x02\x00\x01\x00\x01\xd5\xd3\x00\x11\
379 \x01\x67\x0c\x67\x74\x6c\x64\x2d\x73\x65\x72\x76\x65\x72\x73\
380 \xc0\x42\
381 \x01\x61\xc0\x55\x00\x01\x00\x01\x00\x00\xa3\x1c\
382 \x00\x04\xc0\x05\x06\x1e";
383 let (_, packet) = be_packet(response).unwrap();
384
385 assert_eq!(packet.header.id, 19184);
386 assert_eq!(packet.header.questions_count, 1);
387 assert_eq!(packet.questions[0].qtype, QueryType::A);
388 assert_eq!(packet.questions[0].qclass, QueryClass::IN);
389 assert_eq!(&packet.questions[0].name.to_string()[..], "www.skype.com");
390 assert_eq!(packet.answers.len(), 1);
391 assert_eq!(&packet.answers[0].name.to_string()[..], "www.skype.com");
392 assert_eq!(packet.answers[0].cls, Class::IN);
393 assert_eq!(packet.answers[0].ttl, 3600);
394
395 match &packet.answers[0].data {
396 RData::CName(cname) => {
397 assert_eq!(cname, "livecms.trafficmanager.net");
398 }
399 ref x => panic!("Wrong rdata {x:?}"),
400 }
401 assert_eq!(packet.additional.len(), 1);
402 assert_eq!(
403 &packet.additional[0].name.to_string()[..],
404 "a.gtld-servers.net"
405 );
406 assert_eq!(packet.additional[0].cls, Class::IN);
407 assert_eq!(packet.additional[0].ttl, 41756);
408 match packet.additional[0].data {
409 RData::A(addr) => {
410 assert_eq!(addr, Ipv4Addr::new(192, 5, 6, 30));
411 }
412 ref x => panic!("Wrong rdata {x:?}"),
413 }
414 }
415
416 #[test]
417 fn parse_pack_packet() {
418 let mut response = Packet::default();
419 let address = [
420 IpAddr::V4(Ipv4Addr::new(192, 168, 1, 1)),
421 IpAddr::V6(Ipv6Addr::new(
422 0x2001, 0x0db8, 0x85a3, 0x0000, 0x0000, 0x8a2e, 0x0370, 0x7334,
423 )),
424 ];
425 for ip in address.iter() {
426 let (rtype, rdata) = match ip {
427 IpAddr::V4(ipv4_addr) => (
428 parser::record::Type::A,
429 parser::record::RData::A(*ipv4_addr),
430 ),
431 IpAddr::V6(ipv6_addr) => (
432 parser::record::Type::AAAA,
433 parser::record::RData::AAAA(*ipv6_addr),
434 ),
435 };
436 response.add_answer("example.com", rtype, parser::record::Class::IN, 300, rdata);
437 }
438
439 let srv = Srv::new(0, 0, 6000, "example.com".to_string());
440 response.add_answer(
441 "example.com",
442 parser::record::Type::Srv,
443 parser::record::Class::IN,
444 300,
445 parser::record::RData::Srv(srv),
446 );
447 let packet = response.to_bytes();
448 let (_, parsed_packet) = be_packet(&packet).unwrap();
449 assert_eq!(parsed_packet.header.id, response.header.id);
450 assert_eq!(
451 parsed_packet.header.questions_count,
452 response.header.questions_count
453 );
454 assert_eq!(parsed_packet.answers.len(), response.answers.len());
455 assert_eq!(parsed_packet.nameservers.len(), response.nameservers.len());
456 assert_eq!(parsed_packet.additional.len(), response.additional.len());
457 assert_eq!(parsed_packet.answers[0].name, response.answers[0].name);
458 }
459
460 #[test]
461 fn malformed_packet_does_not_panic() {
462 let packet_hex = "0021641c0000000100000000000078787878787878787878787303636f6d0000100001";
463 let data = decode_hex(packet_hex);
464 let ret = std::panic::catch_unwind(|| {
465 let _ = be_packet(&data);
466 });
467 assert!(ret.is_ok());
468 }
469
470 #[test]
471 fn packet_with_unknown_rr_type_does_not_panic() {
472 let packet_hex = "8116840000010001000000000569627a6c700474657374046d69656b026e6c00000a0001c00c000a0001000000000005497f000001";
473 let data = decode_hex(packet_hex);
474 let ret = std::panic::catch_unwind(|| be_packet(&data));
475 assert!(ret.is_ok());
476 let (_, packet) = be_packet(&data).unwrap();
477 assert_eq!(packet.header.questions_count, 1);
478 assert_eq!(packet.header.answers_count, 1);
479 assert!(packet.answers.is_empty());
480 }
481
482 #[test]
483 fn packet_to_bytes_uses_name_compression_across_sections() {
484 let mut packet = Packet::default();
485 packet.add_question("www.skype.com", QueryType::A, QueryClass::IN, false);
486 packet.add_answer(
487 "mail.skype.com",
488 Type::A,
489 Class::IN,
490 1,
491 RData::A(Ipv4Addr::new(1, 2, 3, 4)),
492 );
493
494 let bytes = packet.to_bytes();
495 let mail_pos = bytes
496 .windows(5)
497 .position(|w| w == b"\x04mail")
498 .expect("mail label missing");
499 let skype_pos = bytes
500 .windows(6)
501 .position(|w| w == b"\x05skype")
502 .expect("skype label missing");
503
504 let ptr = (0xC000u16 | (skype_pos as u16)).to_be_bytes();
505 assert_eq!(&bytes[mail_pos + 5..mail_pos + 7], &ptr);
506 }
507}