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
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}