use crate::{join_multicast, sender_socket, SimpleMdnsError, UNICAST_RESPONSE};
use simple_dns::{rdata::RData, Name, PacketBuf, PacketHeader, Question, QCLASS, QTYPE};
use std::{
net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr, UdpSocket},
time::Duration,
};
pub struct OneShotMdnsResolver {
query_timeout: Duration,
unicast_response: bool,
receiver_socket: UdpSocket,
sender_socket: UdpSocket,
}
impl OneShotMdnsResolver {
pub fn new() -> Result<Self, SimpleMdnsError> {
Ok(Self {
query_timeout: Duration::from_secs(3),
unicast_response: UNICAST_RESPONSE,
receiver_socket: join_multicast(&super::MULTICAST_IPV4_SOCKET)?,
sender_socket: sender_socket(&super::MULTICAST_IPV4_SOCKET)?,
})
}
pub fn query_packet(&self, packet: PacketBuf) -> Result<Option<PacketBuf>, SimpleMdnsError> {
self.sender_socket
.send_to(&packet, &*super::MULTICAST_IPV4_SOCKET)?;
self.get_first_response(packet.packet_id(), self.query_timeout)
}
pub fn query_service_address(
&self,
service_name: &str,
) -> Result<Option<std::net::IpAddr>, SimpleMdnsError> {
let mut packet = PacketBuf::new(PacketHeader::new_query(rand::random(), false), true);
let service_name = Name::new(service_name)?;
packet.add_question(&Question::new(
service_name.clone(),
QTYPE::A,
QCLASS::IN,
self.unicast_response,
))?;
if let Some(response) = self.query_packet(packet)? {
let response = response.to_packet()?;
for anwser in response.answers {
if anwser.name != service_name {
continue;
}
return match anwser.rdata {
RData::A(a) => Ok(Some(IpAddr::V4(Ipv4Addr::from(a.address)))),
RData::AAAA(aaaa) => Ok(Some(IpAddr::V6(Ipv6Addr::from(aaaa.address)))),
_ => Ok(None),
};
}
}
Ok(None)
}
pub fn query_service_address_and_port(
&self,
service_name: &str,
) -> Result<Option<std::net::SocketAddr>, SimpleMdnsError> {
let mut packet = PacketBuf::new(PacketHeader::new_query(rand::random(), false), true);
let parsed_name_service = Name::new(service_name)?;
packet.add_question(&Question::new(
parsed_name_service.clone(),
QTYPE::SRV,
QCLASS::IN,
self.unicast_response,
))?;
if let Some(response) = self.query_packet(packet)? {
let response = response.to_packet()?;
let port = response
.answers
.iter()
.filter(|a| a.name == parsed_name_service && a.match_qtype(QTYPE::SRV))
.find_map(|a| match &a.rdata {
RData::SRV(srv) => Some(srv.port),
_ => None,
});
let mut address = response
.additional_records
.iter()
.filter(|a| a.name == parsed_name_service && a.match_qtype(QTYPE::A))
.find_map(|a| match &a.rdata {
RData::A(a) => Some(IpAddr::V4(Ipv4Addr::from(a.address))),
RData::AAAA(aaaa) => Some(IpAddr::V6(Ipv6Addr::from(aaaa.address))),
_ => None,
});
if port.is_some() && address.is_none() {
address = self.query_service_address(service_name)?;
}
if port.is_some() && address.is_some() {
return Ok(Some(SocketAddr::new(address.unwrap(), port.unwrap())));
}
}
Ok(None)
}
pub fn set_query_timeout(&mut self, query_timeout: Duration) {
self.query_timeout = query_timeout;
}
pub fn set_unicast_response(&mut self, unicast_response: bool) {
self.unicast_response = unicast_response;
}
fn get_first_response(
&self,
packet_id: u16,
query_timeout: Duration,
) -> Result<Option<PacketBuf>, SimpleMdnsError> {
let mut buf = [0u8; 4096];
let timeout = std::time::Instant::now();
loop {
match self.receiver_socket.recv_from(&mut buf[..]) {
Ok((count, _)) => {
if let Ok(header) = PacketHeader::parse(&buf[0..12]) {
if !header.query && header.id == packet_id && header.answers_count > 0 {
return Ok(Some(buf[..count].into()));
}
}
}
Err(_) => {
if timeout.elapsed() > query_timeout {
return Ok(None);
}
}
}
}
}
}
#[cfg(test)]
mod tests {
use std::{str::FromStr, thread};
use crate::{conversion_utils::socket_addr_to_srv_and_address, SimpleMdnsResponder};
use super::*;
fn get_oneshot_responder(srv_name: Name<'static>) -> SimpleMdnsResponder {
let mut responder = SimpleMdnsResponder::default();
let (r1, r2) = socket_addr_to_srv_and_address(
&srv_name,
SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 8080),
0,
);
responder.add_resource(r1);
responder.add_resource(r2);
responder
}
#[test]
fn one_shot_resolver_address_query() {
let _responder = get_oneshot_responder(Name::new_unchecked("_srv._tcp.local"));
thread::sleep(Duration::from_millis(500));
let resolver = OneShotMdnsResolver::new().expect("Failed to create resolver");
let answer = resolver.query_service_address("_srv._tcp.local");
dbg!(&answer);
assert!(answer.is_ok());
let answer = answer.unwrap();
assert!(answer.is_some());
assert_eq!(Ipv4Addr::LOCALHOST, answer.unwrap());
let answer = resolver.query_service_address_and_port("_srv._tcp.local");
dbg!(&answer);
assert!(answer.is_ok());
let answer = answer.unwrap();
assert!(answer.is_some());
assert_eq!(
SocketAddr::from_str("127.0.0.1:8080").unwrap(),
answer.unwrap()
)
}
#[test]
fn one_shot_resolver_timeout() {
let resolver = OneShotMdnsResolver::new().expect("Failed to create resolver");
let answer = resolver.query_service_address("_srv_miss._tcp.local");
dbg!(&answer);
assert!(answer.is_ok());
let answer = answer.unwrap();
assert!(answer.is_none());
}
}