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#[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
36pub 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!(target: "mdns", 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!(target: "mdns", 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
108fn 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 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 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}