1#![allow(clippy::double_parens)]
2use bitfield_struct::bitfield;
3use bytes::BufMut;
4use nom::number::streaming::be_u16;
5
6#[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#[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#[derive(Debug, PartialEq, Eq, Clone, Copy)]
70#[repr(u16)]
71pub enum Opcode {
72 StandardQuery = 0,
74 InverseQuery = 1,
76 ServerStatusRequest = 2,
78 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#[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 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}