Skip to main content

ddns/core/parser/
name.rs

1use std::collections::HashMap;
2
3use bytes::BufMut;
4use nom::{IResult, bytes::streaming::take, number::streaming::be_u8};
5
6pub type Name = String;
7
8pub fn be_name<'a>(input: &'a [u8], origin: &'a [u8]) -> IResult<&'a [u8], Name> {
9    be_name_inner(input, origin, &mut vec![])
10}
11
12/// RFC 1035 4.1.4 的压缩上下文:记录已写入消息中“某个 name 后缀”的偏移量。
13///
14/// key 为后缀域名(例如 `skype.com`),value 为其在消息起始处的偏移(14-bit)。
15#[derive(Debug, Default)]
16pub struct NameCompression {
17    suffix_offsets: HashMap<String, u16>,
18}
19
20impl NameCompression {
21    pub fn new() -> Self {
22        Self::default()
23    }
24
25    fn get_suffix_offset(&self, suffix: &str) -> Option<u16> {
26        self.suffix_offsets.get(suffix).copied()
27    }
28
29    fn remember_suffix(&mut self, suffix: &str, offset: u16) {
30        self.suffix_offsets
31            .entry(suffix.to_string())
32            .or_insert(offset);
33    }
34}
35
36/// 按 RFC 1035 4.1.4 进行压缩编码(若命中已有后缀则写入 pointer 并终止 name)。
37///
38/// 该函数只会引用“已写入到当前消息 earlier 的位置”,并保证指针偏移不超过 14-bit。
39pub fn put_name(buf: &mut Vec<u8>, name: &Name, ctx: &mut NameCompression) -> usize {
40    let start_len = buf.len();
41
42    if name == "." {
43        buf.put_u8(0);
44        return buf.len() - start_len;
45    }
46
47    let trimmed = name.strip_suffix('.').unwrap_or(name);
48    if trimmed.is_empty() {
49        buf.put_u8(0);
50        return buf.len() - start_len;
51    }
52
53    let mut labels = Vec::new();
54    let parts: Vec<&str> = trimmed.split('.').collect();
55    for (i, part) in parts.iter().enumerate() {
56        if part.is_empty() {
57            if i != parts.len() - 1 {
58                tracing::warn!(name, "invalid empty label in middle");
59            }
60            continue;
61        }
62        labels.push(*part);
63    }
64
65    if labels.is_empty() {
66        buf.put_u8(0);
67        return buf.len() - start_len;
68    }
69
70    let mut suffixes = Vec::with_capacity(labels.len());
71    let mut current = String::new();
72    for &label in labels.iter().rev() {
73        if current.is_empty() {
74            current = label.to_string();
75        } else {
76            current = format!("{label}.{current}");
77        }
78        suffixes.push(current.clone());
79    }
80    suffixes.reverse();
81
82    for i in 0..labels.len() {
83        let suffix = &suffixes[i];
84        if let Some(offset) = ctx.get_suffix_offset(suffix) {
85            let ptr = 0xC000u16 | (offset & 0x3FFF);
86            buf.put_u16(ptr);
87            return buf.len() - start_len;
88        }
89
90        if buf.len() <= 0x3FFF {
91            let offset = buf.len() as u16;
92            ctx.remember_suffix(suffix, offset);
93        }
94
95        let label = labels[i];
96        let len = label.len();
97        if len > 63 {
98            tracing::warn!(name, "label exceeds 63 bytes");
99        }
100        buf.put_u8(len as u8);
101        buf.put_slice(label.as_bytes());
102    }
103
104    buf.put_u8(0);
105    buf.len() - start_len
106}
107
108/// 解析一个 DNS name(RFC 1035 3.1 / 4.1.4)。
109///
110/// name 在消息中的编码是 “label 序列 + 0 终止符”,其中每个 label 的格式是:
111/// - `len(1 byte)` + `label bytes(len bytes)`,`len` 取值范围 0..=63
112/// - 当 `len == 0` 表示根标签(root),name 结束
113///
114/// RFC 1035 4.1.4 允许消息压缩:当长度字节的高 2-bit 为 `11`(即 `(len & 0xC0) == 0xC0`)
115/// 时,这两个字节组成一个 14-bit 的偏移量,指向消息起始处的某个 name 后缀:
116/// - `pointer = 0b11xxxxxx xxxxxxxx`,offset 为低 14-bit
117/// - 指针一旦出现,当前 name 立即结束(后面不会再有 root `0`)
118///
119/// 解析时需要防止恶意数据构造 “指针环”,这里用 `visited` 记录递归访问过的 offset,
120/// 遇到重复 offset 即报错。
121fn be_name_inner<'a>(
122    input: &'a [u8],
123    origin: &'a [u8],
124    visited: &mut Vec<usize>,
125) -> IResult<&'a [u8], Name> {
126    let (remain, labels) = be_name_labels(input, origin, visited)?;
127    if labels.is_empty() {
128        return Ok((remain, ".".to_string()));
129    }
130    Ok((remain, labels.join(".")))
131}
132
133fn be_name_labels<'a>(
134    mut input: &'a [u8],
135    origin: &'a [u8],
136    visited: &mut Vec<usize>,
137) -> IResult<&'a [u8], Vec<String>> {
138    let mut labels = Vec::new();
139    loop {
140        let (remain, len) = be_u8(input)?;
141        if len == 0 {
142            return Ok((remain, labels));
143        }
144
145        if (len & 0xC0) == 0xC0 {
146            let (remain, offset_byte) = be_u8(remain)?;
147            let offset = (((len & 0x3F) as u16) << 8) | offset_byte as u16;
148            let offset = offset as usize;
149            if offset >= origin.len() || visited.contains(&offset) {
150                return Err(nom::Err::Error(nom::error::Error::new(
151                    input,
152                    nom::error::ErrorKind::Verify,
153                )));
154            }
155            visited.push(offset);
156            let (_, suffix) = be_name_labels(&origin[offset..], origin, visited)?;
157            visited.pop();
158            labels.extend(suffix);
159            return Ok((remain, labels));
160        }
161
162        if len > 63 {
163            return Err(nom::Err::Error(nom::error::Error::new(
164                input,
165                nom::error::ErrorKind::Verify,
166            )));
167        }
168
169        let (remain, label_bytes) = take(len)(remain)?;
170        labels.push(String::from_utf8_lossy(label_bytes).into_owned());
171        input = remain;
172    }
173}
174
175#[cfg(test)]
176mod test {
177    use super::*;
178
179    fn gen_ascii_label_bytes(len: usize, state: &mut u64) -> Vec<u8> {
180        const ALPHABET: &[u8] = b"abcdefghijklmnopqrstuvwxyz0123456789-";
181        let mut out = Vec::with_capacity(len);
182        for _ in 0..len {
183            *state = state.wrapping_mul(6364136223846793005).wrapping_add(1);
184            let idx = (*state as usize) % ALPHABET.len();
185            out.push(ALPHABET[idx]);
186        }
187        out
188    }
189
190    fn gen_ascii_wire_name(state: &mut u64) -> Vec<u8> {
191        *state = state.wrapping_mul(6364136223846793005).wrapping_add(1);
192        let label_count = 1 + ((*state as usize) % 5);
193        let mut out = Vec::new();
194        for _ in 0..label_count {
195            *state = state.wrapping_mul(6364136223846793005).wrapping_add(1);
196            let len = 1 + ((*state as usize) % 20);
197            out.push(len as u8);
198            out.extend_from_slice(&gen_ascii_label_bytes(len, state));
199        }
200        out.push(0);
201        out
202    }
203
204    #[test]
205    fn parse_example_name() {
206        let name = b"\x07example\x03com\x00";
207        let (remain, parsed_name) = be_name(name, name).unwrap();
208        assert_eq!(remain.len(), 0);
209        assert_eq!(parsed_name, "example.com");
210    }
211
212    #[test]
213    fn parse_badpointer_same_offset() {
214        // A buffer where an offset points to itself,
215        // which is a bad compression pointer.
216        let same_offset = [192, 2, 192, 2];
217        let ret = be_name(&same_offset, &same_offset);
218        assert!(ret.is_err())
219    }
220
221    #[test]
222    fn parse_badpointer_loop_between_offsets() {
223        let buf = [0xC0, 0x02, 0xC0, 0x00];
224        let ret = be_name(&buf, &buf);
225        assert!(ret.is_err());
226    }
227
228    #[test]
229    fn parse_badpointer_out_of_bounds() {
230        let buf = [0xC0, 0x10, 0x00];
231        let ret = be_name(&buf, &buf);
232        assert!(ret.is_err());
233    }
234
235    #[test]
236    fn parse_label_too_long_is_error() {
237        let mut buf = Vec::new();
238        buf.push(64);
239        buf.extend(std::iter::repeat_n(b'a', 64));
240        buf.push(0);
241        let ret = be_name(&buf, &buf);
242        assert!(ret.is_err());
243    }
244
245    #[test]
246    fn parse_pointer_to_root() {
247        let buf = b"\x00\xc0\x00";
248        let (remain, parsed) = be_name(&buf[1..], buf).unwrap();
249        assert!(remain.is_empty());
250        assert_eq!(parsed, ".");
251    }
252
253    #[test]
254    fn pointer_terminates_name_and_does_not_consume_trailing_bytes() {
255        let buf = b"\x03com\x00\x03www\xc0\x00\x00";
256        let (remain, parsed) = be_name(&buf[5..], buf).unwrap();
257        assert_eq!(parsed, "www.com");
258        assert_eq!(remain, b"\x00");
259    }
260
261    #[test]
262    fn parse_chained_compression_pointers() {
263        let buf = b"\x03com\x00\xc0\x00\x03www\xc0\x05";
264        let (remain, parsed) = be_name(&buf[7..], buf).unwrap();
265        assert!(remain.is_empty());
266        assert_eq!(parsed, "www.com");
267    }
268
269    #[test]
270    fn nested_names() {
271        let buf = b"\x02xx\x00\x02yy\xc0\x00\x02zz\xc0\x04";
272
273        let (remaining, parsed) = be_name(buf, buf).unwrap();
274        assert_eq!(remaining.len(), 10);
275        assert_eq!(parsed, "xx");
276
277        let (_remaining, parsed) = be_name(&buf[4..], buf).unwrap();
278        assert_eq!(parsed, "yy.xx");
279
280        // offset only
281        let (_remaining, parsed) = be_name(&buf[7..], buf).unwrap();
282        assert_eq!(parsed, "xx");
283
284        let (_remaining, parsed) = be_name(&buf[9..], buf).unwrap();
285        assert_eq!(parsed, "zz.yy.xx");
286    }
287
288    #[test]
289    fn write_name_compressed_reuses_suffix_pointer() {
290        let mut buf = Vec::new();
291        let mut ctx = NameCompression::new();
292
293        let first = "www.skype.com".to_string();
294        let second = "mail.skype.com".to_string();
295
296        put_name(&mut buf, &first, &mut ctx);
297        let second_pos = buf.len();
298        put_name(&mut buf, &second, &mut ctx);
299
300        assert_eq!(&buf, b"\x03www\x05skype\x03com\x00\x04mail\xc0\x04");
301
302        let (remain, first_parsed) = be_name(&buf, &buf).unwrap();
303        assert_eq!(first_parsed, "www.skype.com");
304        assert_eq!(remain.len(), buf.len() - second_pos);
305
306        let (remain, second_parsed) = be_name(&buf[second_pos..], &buf).unwrap();
307        assert!(remain.is_empty());
308        assert_eq!(second_parsed, "mail.skype.com");
309    }
310
311    #[test]
312    fn write_name_compressed_prefers_longer_suffix() {
313        let mut buf = Vec::new();
314        let mut ctx = NameCompression::new();
315
316        let first = "a.b.c.com".to_string();
317        let second = "x.c.com".to_string();
318
319        put_name(&mut buf, &first, &mut ctx);
320        put_name(&mut buf, &second, &mut ctx);
321
322        assert_eq!(&buf, b"\x01a\x01b\x01c\x03com\x00\x01x\xc0\x04");
323    }
324
325    #[test]
326    fn write_and_parse_roundtrip() {
327        let raw_name = "HP Color LaserJet Pro M478f-9f [EC3C83]._http._tcp.local".to_string();
328
329        let mut buf = Vec::new();
330        let mut ctx = NameCompression::new();
331        let _ = put_name(&mut buf, &raw_name, &mut ctx);
332
333        let (remaining, parsed) = be_name(&buf, &buf).unwrap();
334        assert!(remaining.is_empty());
335        assert_eq!(parsed, raw_name);
336    }
337
338    #[test]
339    fn random_ascii_name_roundtrip_without_compression_pointer() {
340        let mut state = 0x1234_5678_9abc_def0u64;
341        for _ in 0..1000 {
342            let wire = gen_ascii_wire_name(&mut state);
343            let (_, name) = be_name(&wire, &wire).unwrap();
344            let mut buf = Vec::new();
345            let mut ctx = NameCompression::new();
346            put_name(&mut buf, &name, &mut ctx);
347            assert_eq!(buf, wire);
348        }
349    }
350}