Skip to main content

ddns/core/parser/
header.rs

1#![allow(clippy::double_parens)]
2use bitfield_struct::bitfield;
3use bytes::BufMut;
4use nom::number::streaming::be_u16;
5
6/// See <https://datatracker.ietf.org/doc/html/rfc1035#autoid-40>
7/// 与标准 DNS 不同,flags 字段的 zero 3bits 后两 bits 用于 AD、CD
8/// See <https://datatracker.ietf.org/doc/html/rfc6762#autoid-48>
9/// ```text
10/// 0  1  2  3  4  5  6  7  8  9  0  1  2  3  4  5
11/// +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+
12/// |                      ID                       |
13/// +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+
14/// |QR|   Opcode  |AA|TC|RD|RA|Z |AD|CD|   RCODE   |
15/// +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+
16/// |                    QDCOUNT                    |
17/// +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+
18/// |                    ANCOUNT                    |
19/// +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+
20/// |                    NSCOUNT                    |
21/// +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+
22/// |                    ARCOUNT                    |
23/// +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+
24/// ```
25#[derive(Debug, Clone)]
26pub struct Header {
27    pub(crate) id: u16,
28    pub(crate) flags: Flags,
29    pub(crate) questions_count: u16,
30    pub(crate) answers_count: u16,
31    pub(crate) nameservers_count: u16,
32    pub(crate) additional_count: u16,
33}
34
35impl Default for Header {
36    fn default() -> Self {
37        let flags = Flags::new();
38        flags.with_query(false).set_opcode(Opcode::StandardQuery);
39        Self {
40            id: 0,
41            flags,
42            questions_count: 0,
43            answers_count: 0,
44            nameservers_count: 0,
45            additional_count: 0,
46        }
47    }
48}
49
50/// See <https://datatracker.ietf.org/doc/html/rfc6762#autoid-48>
51#[bitfield(u16, order = Msb)]
52#[derive(PartialEq, Eq)]
53pub struct Flags {
54    pub(crate) query: bool,
55    #[bits(4)]
56    pub(crate) opcode: Opcode,
57    pub(crate) authoritative: bool,
58    pub(crate) trun_cache: bool,
59    pub(crate) recursion_desired: bool,
60    pub(crate) recursion_available: bool,
61    pub(crate) zero: bool,
62    pub(crate) authenticated_data: bool,
63    pub(crate) checking_disabled: bool,
64    #[bits(4)]
65    pub(crate) response_code: ResponseCode,
66}
67
68/// The OPCODE value according to RFC 1035
69#[derive(Debug, PartialEq, Eq, Clone, Copy)]
70#[repr(u16)]
71pub enum Opcode {
72    /// Normal query
73    StandardQuery = 0,
74    /// Inverse query (query a name by IP)
75    InverseQuery = 1,
76    /// Server status request
77    ServerStatusRequest = 2,
78    /// Reserved opcode for future use
79    Reserved(u8),
80}
81
82impl Opcode {
83    const fn into_bits(self) -> u8 {
84        match self {
85            Self::StandardQuery => 0,
86            Self::InverseQuery => 1,
87            Self::ServerStatusRequest => 2,
88            Self::Reserved(value) => value,
89        }
90    }
91
92    const fn from_bits(value: u8) -> Self {
93        match value {
94            0 => Self::StandardQuery,
95            1 => Self::InverseQuery,
96            2 => Self::ServerStatusRequest,
97            _ => Self::Reserved(value),
98        }
99    }
100}
101
102// The RCODE value according to RFC 1035
103#[derive(Debug, PartialEq, Eq, Clone, Copy)]
104#[repr(u16)]
105pub enum ResponseCode {
106    NoError = 0,
107    FormatError = 1,
108    ServerFailure = 2,
109    NameError = 3,
110    NotImplemented = 4,
111    Refused = 5,
112    Reserved(u8),
113}
114
115impl ResponseCode {
116    const fn into_bits(self) -> u8 {
117        match self {
118            Self::NoError => 0,
119            Self::FormatError => 1,
120            Self::ServerFailure => 2,
121            Self::NameError => 3,
122            Self::NotImplemented => 4,
123            Self::Refused => 5,
124            Self::Reserved(value) => value,
125        }
126    }
127
128    const fn from_bits(value: u8) -> Self {
129        match value {
130            0 => Self::NoError,
131            1 => Self::FormatError,
132            2 => Self::ServerFailure,
133            3 => Self::NameError,
134            4 => Self::NotImplemented,
135            5 => Self::Refused,
136            _ => Self::Reserved(value),
137        }
138    }
139}
140
141pub fn be_header(input: &[u8]) -> nom::IResult<&[u8], Header> {
142    let (remain, id) = be_u16(input)?;
143    let (remain, flags) = be_u16(remain)?;
144    let flags = Flags::from(flags);
145    let (remain, questions) = be_u16(remain)?;
146    let (remain, answers) = be_u16(remain)?;
147    let (remain, nameservers) = be_u16(remain)?;
148    let (remain, additional) = be_u16(remain)?;
149    Ok((
150        remain,
151        Header {
152            id,
153            flags,
154            questions_count: questions,
155            answers_count: answers,
156            nameservers_count: nameservers,
157            additional_count: additional,
158        },
159    ))
160}
161
162pub trait WriteHeader {
163    fn put_header(&mut self, header: &Header);
164}
165
166impl<T: BufMut> WriteHeader for T {
167    fn put_header(&mut self, header: &Header) {
168        self.put_u16(header.id);
169        self.put_u16(header.flags.into());
170        self.put_u16(header.questions_count);
171        self.put_u16(header.answers_count);
172        self.put_u16(header.nameservers_count);
173        self.put_u16(header.additional_count);
174    }
175}
176
177#[cfg(test)]
178mod test {
179    use bytes::BytesMut;
180    use nom::AsBytes;
181
182    use super::*;
183
184    #[test]
185    fn parse_example_query() {
186        let query = b"\x06%\x01\x00\x00\x01\x00\x00\x00\x00\x00\x00\
187                      \x07example\x03com\x00\x00\x01\x00\x01";
188
189        let (remain, header) = be_header(query).unwrap();
190        assert_eq!(remain.len(), query.len() - 12);
191        assert_eq!(header.id, 1573);
192        assert_eq!(header.questions_count, 1);
193        assert_eq!(header.answers_count, 0);
194        assert_eq!(header.nameservers_count, 0);
195        assert_eq!(header.additional_count, 0);
196        let flags = Flags::new()
197            .with_recursion_desired(true)
198            .with_response_code(ResponseCode::NoError)
199            .with_opcode(Opcode::StandardQuery);
200        assert_eq!(header.flags.into_bits(), 0x0100);
201        assert_eq!(header.flags, flags);
202    }
203
204    #[test]
205    fn parse_example_response() {
206        let response = b"\x06%\x81\x80\x00\x01\x00\x01\x00\x00\x00\x00\
207                     \x07example\x03com\x00\x00\x01\x00\x01\
208                     \xc0\x0c\x00\x01\x00\x01\x00\x00\x04\xf8\
209                     \x00\x04]\xb8\xd8\"";
210
211        let (remain, header) = be_header(response).unwrap();
212        assert_eq!(remain.len(), response.len() - 12);
213        assert_eq!(header.id, 1573);
214        assert_eq!(header.questions_count, 1);
215        assert_eq!(header.answers_count, 1);
216        assert_eq!(header.nameservers_count, 0);
217        assert_eq!(header.additional_count, 0);
218        // response
219        let flag = Flags::new()
220            .with_recursion_desired(true)
221            .with_recursion_available(true)
222            .with_query(true)
223            .with_response_code(ResponseCode::NoError)
224            .with_opcode(Opcode::StandardQuery);
225        assert_eq!(header.flags.into_bits(), 0x8180);
226        assert_eq!(header.flags, flag);
227    }
228
229    #[test]
230    fn parse_query_with_ad_set() {
231        let query = b"\x06%\x01\x20\x00\x01\x00\x00\x00\x00\x00\x00\
232                  \x07example\x03com\x00\x00\x01\x00\x01";
233        let (remain, header) = be_header(query).unwrap();
234        assert_eq!(remain.len(), query.len() - 12);
235        assert_eq!(header.id, 1573);
236        assert_eq!(header.questions_count, 1);
237        assert_eq!(header.answers_count, 0);
238        assert_eq!(header.nameservers_count, 0);
239        assert_eq!(header.additional_count, 0);
240
241        let flags = Flags::new()
242            .with_recursion_desired(true)
243            .with_authenticated_data(true)
244            .with_response_code(ResponseCode::NoError)
245            .with_opcode(Opcode::StandardQuery);
246        assert_eq!(header.flags.into_bits(), 0x0120);
247        assert_eq!(header.flags, flags);
248    }
249
250    #[test]
251    fn parse_query_with_cd_set() {
252        let query = b"\x06%\x01\x10\x00\x01\x00\x00\x00\x00\x00\x00\
253                      \x07example\x03com\x00\x00\x01\x00\x01";
254        let (remain, header) = be_header(query).unwrap();
255        assert_eq!(remain.len(), query.len() - 12);
256        assert_eq!(header.id, 1573);
257        assert_eq!(header.questions_count, 1);
258        assert_eq!(header.answers_count, 0);
259        assert_eq!(header.nameservers_count, 0);
260        assert_eq!(header.additional_count, 0);
261        let flags = Flags::new()
262            .with_recursion_desired(true)
263            .with_checking_disabled(true)
264            .with_response_code(ResponseCode::NoError)
265            .with_opcode(Opcode::StandardQuery);
266        assert_eq!(header.flags.into_bits(), 0x0110);
267        assert_eq!(header.flags, flags);
268    }
269
270    #[test]
271    fn write_example_query() {
272        let header = Header {
273            id: 1573,
274            flags: Flags::from(0x0100),
275            questions_count: 1,
276            answers_count: 0,
277            nameservers_count: 0,
278            additional_count: 0,
279        };
280        let mut buf = BytesMut::with_capacity(12);
281        buf.put_header(&header);
282        assert_eq!(
283            buf.as_bytes(),
284            b"\x06%\x01\x00\x00\x01\x00\x00\x00\x00\x00\x00"
285        );
286    }
287
288    #[test]
289    fn write_example_response() {
290        let header = Header {
291            id: 1573,
292            flags: Flags::from(0x8180),
293            questions_count: 1,
294            answers_count: 1,
295            nameservers_count: 0,
296            additional_count: 0,
297        };
298        let mut buf = BytesMut::with_capacity(28);
299        buf.put_header(&header);
300        assert_eq!(
301            buf.as_bytes(),
302            b"\x06%\x81\x80\x00\x01\x00\x01\x00\x00\x00\x00"
303        );
304    }
305}