use std::net::Ipv4Addr;
pub fn build_mdns_packet(name: &str, domain: &str, ip: Ipv4Addr) -> Vec<u8> {
build_mdns_packet_with_ttl(name, domain, ip, 120)
}
pub fn build_goodbye_packet(name: &str, domain: &str, ip: Ipv4Addr) -> Vec<u8> {
build_mdns_packet_with_ttl(name, domain, ip, 0)
}
fn build_mdns_packet_with_ttl(name: &str, domain: &str, ip: Ipv4Addr, ttl: u32) -> Vec<u8> {
let mut packet = Vec::new();
packet.extend_from_slice(&[0x00, 0x00]); packet.extend_from_slice(&[0x84, 0x00]); packet.extend_from_slice(&[0x00, 0x00]); packet.extend_from_slice(&[0x00, 0x01]); packet.extend_from_slice(&[0x00, 0x00]); packet.extend_from_slice(&[0x00, 0x00]);
let full_name = format!("{}.{}", name, domain);
for label in full_name.split('.') {
packet.push(label.len() as u8);
packet.extend_from_slice(label.as_bytes());
}
packet.push(0x00);
packet.extend_from_slice(&[0x00, 0x01]);
packet.extend_from_slice(&[0x80, 0x01]);
packet.extend_from_slice(&ttl.to_be_bytes());
packet.extend_from_slice(&[0x00, 0x04]);
packet.extend_from_slice(&ip.octets());
packet
}
pub fn build_mdns_query(name: &str, domain: &str) -> Vec<u8> {
let mut packet = Vec::new();
packet.extend_from_slice(&[0x00, 0x00]); packet.extend_from_slice(&[0x00, 0x00]); packet.extend_from_slice(&[0x00, 0x01]); packet.extend_from_slice(&[0x00, 0x00]); packet.extend_from_slice(&[0x00, 0x00]); packet.extend_from_slice(&[0x00, 0x00]);
let full_name = format!("{}.{}", name, domain);
for label in full_name.split('.') {
packet.push(label.len() as u8);
packet.extend_from_slice(label.as_bytes());
}
packet.push(0x00);
packet.extend_from_slice(&[0x00, 0x01]);
packet.extend_from_slice(&[0x00, 0x01]);
packet
}
pub fn parse_mdns_query(packet: &[u8]) -> Option<String> {
if packet.len() < 12 {
return None;
}
if packet[2] & 0x80 != 0 {
return None;
}
let question_count = u16::from_be_bytes([packet[4], packet[5]]);
if question_count == 0 {
return None;
}
let mut pos = 12;
let mut name_parts = Vec::new();
while pos < packet.len() {
let len = packet[pos] as usize;
if len == 0 {
break;
}
pos += 1;
if pos + len > packet.len() {
return None;
}
if let Ok(label) = std::str::from_utf8(&packet[pos..pos + len]) {
name_parts.push(label.to_string());
} else {
return None;
}
pos += len;
}
if name_parts.is_empty() {
None
} else {
Some(name_parts.join("."))
}
}
pub fn parse_mdns_response(packet: &[u8]) -> Option<(String, Ipv4Addr)> {
if packet.len() < 12 {
return None;
}
if packet[2] & 0x80 == 0 {
return None;
}
let answer_count = u16::from_be_bytes([packet[6], packet[7]]);
if answer_count == 0 {
return None;
}
let mut pos = 12;
let question_count = u16::from_be_bytes([packet[4], packet[5]]);
for _ in 0..question_count {
while pos < packet.len() && packet[pos] != 0 {
let len = packet[pos] as usize;
pos += 1 + len;
}
pos += 1;
pos += 4; }
if pos >= packet.len() {
return None;
}
let mut name_parts = Vec::new();
while pos < packet.len() {
let len = packet[pos] as usize;
if len >= 0xC0 {
if pos + 1 >= packet.len() {
return None;
}
let offset = (u16::from_be_bytes([packet[pos] & 0x3F, packet[pos + 1]])) as usize;
let mut offset_pos = offset;
while offset_pos < packet.len() {
let offset_len = packet[offset_pos] as usize;
if offset_len == 0 || offset_len >= 0xC0 {
break;
}
offset_pos += 1;
if offset_pos + offset_len > packet.len() {
return None;
}
if let Ok(label) = std::str::from_utf8(&packet[offset_pos..offset_pos + offset_len]) {
name_parts.push(label.to_string());
}
offset_pos += offset_len;
}
pos += 2; break;
}
if len == 0 {
pos += 1; break;
}
pos += 1;
if pos + len > packet.len() {
return None;
}
if let Ok(label) = std::str::from_utf8(&packet[pos..pos + len]) {
name_parts.push(label.to_string());
} else {
return None;
}
pos += len;
}
if name_parts.is_empty() {
return None;
}
if pos + 10 > packet.len() {
return None;
}
let rtype = u16::from_be_bytes([packet[pos], packet[pos + 1]]);
pos += 2;
pos += 2;
pos += 4;
let rdlength = u16::from_be_bytes([packet[pos], packet[pos + 1]]);
pos += 2;
if rtype != 1 || rdlength != 4 {
return None;
}
if pos + 4 > packet.len() {
return None;
}
let ip = Ipv4Addr::new(packet[pos], packet[pos + 1], packet[pos + 2], packet[pos + 3]);
Some((name_parts.join("."), ip))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_build_mdns_packet() {
let ip = Ipv4Addr::new(192, 168, 1, 50);
let packet = build_mdns_packet("testhost", "local", ip);
assert_eq!(packet[2], 0x84); assert_eq!(packet[7], 0x01);
assert!(packet.len() > 12);
let ip_octets = ip.octets();
let packet_len = packet.len();
assert_eq!(&packet[packet_len - 4..], &ip_octets);
}
#[test]
fn test_parse_mdns_query() {
let mut packet = vec![
0x00, 0x00, 0x00, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, ];
packet.push(8); packet.extend_from_slice(b"testhost");
packet.push(5); packet.extend_from_slice(b"local");
packet.push(0);
let parsed = parse_mdns_query(&packet);
assert_eq!(parsed, Some("testhost.local".to_string()));
}
#[test]
fn test_parse_mdns_response() {
let packet = vec![
0x00, 0x00, 0x84, 0x00, 0x00, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00, 0x00, ];
let parsed = parse_mdns_query(&packet);
assert_eq!(parsed, None);
}
#[test]
fn test_build_mdns_query() {
let packet = build_mdns_query("testhost", "local");
assert_eq!(packet[2], 0x00); assert_eq!(packet[3], 0x00); assert_eq!(packet[5], 0x01);
assert!(packet.len() > 12);
}
#[test]
fn test_build_goodbye_packet() {
let ip = Ipv4Addr::new(192, 168, 1, 50);
let packet = build_goodbye_packet("testhost", "local", ip);
assert_eq!(packet[2], 0x84); assert_eq!(packet[7], 0x01);
let ttl_pos = 32;
let ttl = u32::from_be_bytes([
packet[ttl_pos],
packet[ttl_pos + 1],
packet[ttl_pos + 2],
packet[ttl_pos + 3],
]);
assert_eq!(ttl, 0);
let ip_octets = ip.octets();
let packet_len = packet.len();
assert_eq!(&packet[packet_len - 4..], &ip_octets);
}
}