Skip to main content

ddns/core/parser/
packet.rs

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/// Parsed DNS packet
18#[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        // true 代表是 response
168        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
242/// Minimum possible wire length of one DNS question.
243const MIN_QUESTION_WIRE_LEN: usize = 5;
244/// Minimum possible wire length of one DNS resource record.
245const MIN_RESOURCE_RECORD_WIRE_LEN: usize = 11;
246
247/// Verify that the input can possibly hold every section declared by the header.
248fn 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, // ID
509            0, 0, // flags
510            0xff, 0xff, // questions
511            0, 0, // answers
512            0, 0, // nameservers
513            0, 0, // additional
514        ];
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, // ID
523            0, 0, // flags
524            0, 0, // questions
525            0, 1, // answers
526            0, 1, // nameservers
527            0, 1, // additional
528        ];
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}