1use bstr::{BStr, ByteSlice};
2use hickory_proto::rr::Name;
3use nom::branch::alt;
4use nom::bytes::complete::{take_while1, take_while_m_n};
5use nom::combinator::{map_res, opt, recognize};
6use nom::error::{context, ContextError, ErrorKind, FromExternalError, ParseError as _};
7use nom::multi::{many0, many1};
8use nom::sequence::pair;
9use nom::{Input, Parser as _};
10use nom_locate::LocatedSpan;
11use std::fmt::{self, Debug, Write};
12use std::hash::Hash;
13use std::marker::PhantomData;
14use std::net::{Ipv4Addr, Ipv6Addr};
15use std::str::FromStr;
16
17pub type Span<'a> = LocatedSpan<&'a [u8]>;
18pub type IResult<'a, A, B> = nom::IResult<A, B, ParseError<Span<'a>>>;
19
20pub fn make_span(s: &'_ [u8]) -> Span<'_> {
21 Span::new(s)
22}
23
24pub fn tag<E>(tag: &'static str) -> TagParser<E> {
28 TagParser {
29 tag,
30 no_case: false,
31 e: PhantomData,
32 }
33}
34
35pub fn tag_no_case<E>(tag: &'static str) -> TagParser<E> {
36 TagParser {
37 tag,
38 no_case: true,
39 e: PhantomData,
40 }
41}
42
43pub struct TagParser<E> {
45 tag: &'static str,
46 no_case: bool,
47 e: PhantomData<E>,
48}
49
50impl<I, Error: nom::error::ParseError<I> + nom::error::FromExternalError<I, String>> nom::Parser<I>
52 for TagParser<Error>
53where
54 I: nom::Input + nom::Compare<&'static str> + nom::AsBytes,
55{
56 type Output = I;
57 type Error = Error;
58
59 fn process<OM: nom::OutputMode>(
60 &mut self,
61 i: I,
62 ) -> nom::PResult<OM, I, Self::Output, Self::Error> {
63 use nom::error::ErrorKind;
64 use nom::{CompareResult, Err, Mode};
65
66 let tag_len = self.tag.input_len();
67
68 let compare_result = if self.no_case {
69 i.compare_no_case(self.tag)
70 } else {
71 i.compare(self.tag)
72 };
73
74 match compare_result {
75 CompareResult::Ok => Ok((i.take_from(tag_len), OM::Output::bind(|| i.take(tag_len)))),
76 CompareResult::Incomplete => Err(Err::Error(OM::Error::bind(|| {
77 Error::from_external_error(
78 i,
79 ErrorKind::Fail,
80 format!(
81 "expected \"{}\" but ran out of input",
82 self.tag.escape_debug()
83 ),
84 )
85 }))),
86
87 CompareResult::Error => {
88 let available = i.take(i.input_len().min(tag_len));
89 Err(Err::Error(OM::Error::bind(|| {
90 Error::from_external_error(
91 i,
92 ErrorKind::Fail,
93 format!(
94 "expected \"{}\" but found {:?}",
95 self.tag.escape_debug(),
96 BStr::new(available.as_bytes())
97 ),
98 )
99 })))
100 }
101 }
102 }
103}
104
105#[derive(Debug)]
106pub enum ParseErrorKind {
107 Context(&'static str),
108 Char(char),
109 Nom(ErrorKind),
110 External { kind: ErrorKind, reason: String },
111}
112
113#[derive(Debug)]
114pub struct ParseError<I: Debug> {
115 pub errors: Vec<(I, ParseErrorKind)>,
116}
117
118impl<I: Debug> ContextError<I> for ParseError<I> {
119 fn add_context(input: I, ctx: &'static str, mut other: Self) -> Self {
120 other.errors.push((input, ParseErrorKind::Context(ctx)));
121 other
122 }
123}
124
125impl<I: Debug> nom::error::ParseError<I> for ParseError<I> {
126 fn from_error_kind(input: I, kind: ErrorKind) -> Self {
127 Self {
128 errors: vec![(input, ParseErrorKind::Nom(kind))],
129 }
130 }
131
132 fn append(input: I, kind: ErrorKind, mut other: Self) -> Self {
133 other.errors.push((input, ParseErrorKind::Nom(kind)));
134 other
135 }
136
137 fn from_char(input: I, c: char) -> Self {
138 Self {
139 errors: vec![(input, ParseErrorKind::Char(c))],
140 }
141 }
142}
143
144impl<I: Debug, E: std::fmt::Display> nom::error::FromExternalError<I, E> for ParseError<I> {
145 fn from_external_error(input: I, kind: ErrorKind, err: E) -> Self {
146 Self {
147 errors: vec![(
148 input,
149 ParseErrorKind::External {
150 kind,
151 reason: format!("{err:#}"),
152 },
153 )],
154 }
155 }
156}
157
158pub fn make_context_error<S: Into<String>>(
159 input: Span<'_>,
160 reason: S,
161) -> nom::Err<ParseError<Span<'_>>> {
162 nom::Err::Error(ParseError {
163 errors: vec![(
164 input,
165 ParseErrorKind::External {
166 kind: nom::error::ErrorKind::Fail,
167 reason: reason.into(),
168 },
169 )],
170 })
171}
172
173pub fn explain_nom(input: Span, err: nom::Err<ParseError<Span<'_>>>) -> String {
174 match err {
175 nom::Err::Error(e) => {
176 let mut result = String::new();
177 let mut lines_shown = vec![];
178
179 for (span, kind) in e.errors.iter() {
180 if input.is_empty() {
181 match kind {
182 ParseErrorKind::Char(c) => {
183 write!(&mut result, "Error expected '{c}', got empty input\n\n")
184 }
185 ParseErrorKind::Context(s) => {
186 write!(&mut result, "Error in {s}, got empty input\n\n")
187 }
188 ParseErrorKind::External { kind, reason } => {
189 write!(&mut result, "Error {reason} {kind:?}, got empty input\n\n")
190 }
191 ParseErrorKind::Nom(e) => {
192 write!(&mut result, "Error in {e:?}, got empty input\n\n")
193 }
194 }
195 .ok();
196 continue;
197 }
198
199 let line_number = span.location_line();
200 let input_line = span.get_line_beginning();
201 let mut line = String::new();
205 for (start, end, c) in input_line.char_indices() {
206 let c = match c {
207 '\t' => '\u{2409}',
208 '\r' => '\u{240d}',
209 '\n' => '\u{240a}',
210 c => c,
211 };
212
213 if c == std::char::REPLACEMENT_CHARACTER {
214 let bytes = &input_line[start..end];
215 for b in bytes.iter() {
216 line.push_str(&format!("\\x{b:02X}"));
217 }
218 } else {
219 line.push(c);
220 }
221 }
222
223 let column = span.get_utf8_column();
224
225 lines_shown.push(line_number);
226
227 let mut caret = " ".repeat(column.saturating_sub(1));
228 caret.push('^');
229 for _ in 1..span.fragment().len() {
230 caret.push('_')
231 }
232
233 match kind {
234 ParseErrorKind::Char(expected) => {
235 if let Some(actual) = span.fragment().chars().next() {
236 write!(
237 &mut result,
238 "Error at line {line_number}:\n\
239 {line}\n\
240 {caret}\n\
241 expected '{expected}', found {actual}\n\n",
242 )
243 } else {
244 write!(
245 &mut result,
246 "Error at line {line_number}:\n\
247 {line}\n\
248 {caret}\n\
249 expected '{expected}', got end of input\n\n",
250 )
251 }
252 }
253 ParseErrorKind::Context(context) => {
254 write!(&mut result, "while parsing {context}\n")
255 }
256 ParseErrorKind::External { kind: _, reason } => {
257 write!(
258 &mut result,
259 "Error at line {line_number}, {reason}:\n\
260 {line}\n\
261 {caret}\n\n",
262 )
263 }
264 ParseErrorKind::Nom(nom_err) => {
265 write!(
266 &mut result,
267 "Error at line {line_number}, in {nom_err:?}:\n\
268 {line}\n\
269 {caret}\n\n",
270 )
271 }
272 }
273 .ok();
274 }
275 result
276 }
277 _ => format!("{err:#}"),
278 }
279}
280
281pub fn utf8_non_ascii(input: Span) -> IResult<Span, Span> {
292 use nom::Err;
293
294 match input.char_indices().next() {
295 Some((start, end, c)) => {
296 let len = end - start;
297 if c as u32 <= 0x7f {
298 return Err(Err::Error(ParseError::from_error_kind(
300 input,
301 ErrorKind::Fail,
302 )));
303 }
304 let slice = &input[start..end];
305 if c == std::char::REPLACEMENT_CHARACTER {
306 let mut verify = [0u8; 4];
307 if slice != c.encode_utf8(&mut verify).as_bytes() {
308 return Err(Err::Error(ParseError::from_error_kind(
311 input,
312 ErrorKind::Fail,
313 )));
314 }
315 }
316 Ok((input.take_from(len), input.take(len)))
318 }
319 None => {
320 Err(Err::Error(ParseError::from_error_kind(
322 input,
323 ErrorKind::Eof,
324 )))
325 }
326 }
327}
328
329fn snum(input: Span) -> IResult<Span, Span> {
330 take_while_m_n(1, 3, |c: u8| c.is_ascii_digit()).parse(input)
331}
332
333pub fn ipv4_address(input: Span) -> IResult<Span, Ipv4Addr> {
334 context(
335 "ipv4_address",
336 map_res(
337 recognize((snum, tag("."), snum, tag("."), snum, tag("."), snum)),
338 |matched| {
339 let v4str = std::str::from_utf8(&matched).expect("can only be ascii");
340 v4str.parse().map_err(|err| {
341 nom::Err::Error(ParseError::from_external_error(
342 input,
343 ErrorKind::Fail,
344 format!("invalid ipv4_address: {err}"),
345 ))
346 })
347 },
348 ),
349 )
350 .parse(input)
351}
352
353pub fn ipv6_address(input: Span) -> IResult<Span, Ipv6Addr> {
354 context(
355 "ipv6_address",
356 map_res(
357 take_while1(|c: u8| c.is_ascii_hexdigit() || c == b':' || c == b'.'),
358 |matched: Span| {
359 let v6str = std::str::from_utf8(&matched).expect("can only be ascii");
360 v6str.parse().map_err(|err| {
361 nom::Err::Error(ParseError::from_external_error(
362 input,
363 ErrorKind::Fail,
364 format!("invalid ipv6_address: {err}"),
365 ))
366 })
367 },
368 ),
369 )
370 .parse(input)
371}
372
373#[derive(Clone, Debug, PartialEq, Eq, Hash)]
377pub struct DomainString(String);
378
379impl fmt::Display for DomainString {
380 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
381 f.write_str(&self.0)
382 }
383}
384
385impl DomainString {
386 pub fn name(&self) -> Name {
387 Name::from_str_relaxed(&self.0)
388 .expect("cannot construct DomainString with an invalid domain name")
389 }
390
391 pub fn as_str(&self) -> &str {
393 &self.0
394 }
395}
396
397impl FromStr for DomainString {
398 type Err = String;
399 fn from_str(s: &str) -> Result<Self, Self::Err> {
400 let name = Name::from_str_relaxed(s).map_err(|e| e.to_string())?;
401 Ok(Self(name.to_ascii()))
402 }
403}
404
405impl From<DomainString> for Name {
406 fn from(val: DomainString) -> Self {
407 val.name()
408 }
409}
410
411impl From<&DomainString> for Name {
412 fn from(val: &DomainString) -> Self {
413 val.name()
414 }
415}
416
417fn let_dig(input: Span) -> IResult<Span, Span> {
419 recognize(alt((
420 take_while_m_n(1, 1, |c: u8| c.is_ascii_alphanumeric()),
421 utf8_non_ascii,
422 )))
423 .parse(input)
424}
425
426fn ldh_str(input: Span) -> IResult<Span, Span> {
432 recognize(many1(alt((
433 take_while_m_n(1, 1, |c: u8| {
434 c.is_ascii_alphanumeric() || c == b'-' || c == b'_'
435 }),
436 utf8_non_ascii,
437 ))))
438 .parse(input)
439}
440
441fn sub_domain(input: Span) -> IResult<Span, Span> {
443 recognize(pair(let_dig, opt(ldh_str))).parse(input)
444}
445
446pub fn domain_name(input: Span) -> IResult<Span, DomainString> {
448 context(
449 "domain-name",
450 map_res(
451 recognize(pair(sub_domain, many0(pair(tag("."), sub_domain)))),
452 |matched: Span| match std::str::from_utf8(&matched) {
453 Ok(s) => s.parse().map_err(|err| {
454 nom::Err::Error(ParseError::from_external_error(
455 input,
456 ErrorKind::Fail,
457 format!("invalid domain name: {err}"),
458 ))
459 }),
460 Err(err) => Err(nom::Err::Error(ParseError::from_external_error(
461 input,
462 ErrorKind::Fail,
463 format!("invalid domain name: {err}"),
464 ))),
465 },
466 ),
467 )
468 .parse(input)
469}
470
471#[cfg(test)]
472mod tests {
473 use super::*;
474
475 #[test]
476 fn test_ipv4_parse() {
477 let (_, addr) = ipv4_address(make_span(b"192.168.1.1")).unwrap();
479 k9::assert_equal!(addr, Ipv4Addr::new(192, 168, 1, 1));
480 }
481
482 #[test]
483 fn test_ipv6_parse() {
484 let (_, v6a) = ipv6_address(make_span(b"2001:0db8:0000:0000:0000:0000:0000:0001")).unwrap();
487 let (_, v6b) = ipv6_address(make_span(b"2001:db8::1")).unwrap();
488 k9::assert_equal!(v6a, v6b);
489 }
490
491 #[test]
492 fn test_domain_string_partial_eq() {
493 let d1 = DomainString::from_str("EXAMPLE.COM").unwrap();
495 let d2 = DomainString::from_str("example.com").unwrap();
496
497 assert_eq!(d1, d2);
498 }
499
500 #[test]
501 fn test_domain_string_partial_eq_idna() {
502 let d1 = DomainString::from_str("münchen.de").unwrap();
504 let d2 = DomainString::from_str("xn--mnchen-3ya.de").unwrap();
505
506 assert_eq!(d1, d2);
507 }
508}