Skip to main content

simple_mdns/sync_discovery/
oneshot_resolver.rs

1use crate::{
2    socket_helper::{join_multicast, sender_socket},
3    NetworkScope, SimpleMdnsError, UNICAST_RESPONSE,
4};
5use simple_dns::{header_buffer, rdata::RData, Name, Packet, Question, CLASS, TYPE};
6
7use std::{
8    net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr, UdpSocket},
9    time::{Duration, Instant},
10};
11
12/// Provides One Shot queries (legacy mDNS)
13///
14/// Every query will timeout after `query_timeout` elapses (defaults to 3 seconds)
15///
16/// One Shot queries returns only the first valid response to arrive
17/// ```
18///     use simple_mdns::sync_discovery::OneShotMdnsResolver;
19///     use std::time::Duration;
20///     
21///     let mut resolver = OneShotMdnsResolver::new().expect("Can't create one shot resolver");
22///     resolver.set_query_timeout(Duration::from_secs(1));
23///     
24///     // querying for IP Address
25///     let answer = resolver.query_service_address("_myservice._tcp.local").expect("Failed to query service address");
26///     println!("{:?}", answer);
27///     // IpV4Addr or IpV6Addr, depending on what was returned
28///    
29///     let answer = resolver.query_service_address_and_port("_myservice._tcp.local").expect("Failed to query service address and port");
30///     println!("{:?}", answer);
31///     // SocketAddr, "127.0.0.1:8080", with a ipv4 or ipv6
32/// ```
33#[derive(Debug)]
34pub struct OneShotMdnsResolver {
35    query_timeout: Duration,
36    unicast_response: bool,
37    receiver_socket: UdpSocket,
38    sender_socket: UdpSocket,
39    network_scope: NetworkScope,
40}
41
42impl OneShotMdnsResolver {
43    /// Creates a new OneShotMdnsResolver using IP V4 with unspecified interface
44    pub fn new() -> Result<Self, SimpleMdnsError> {
45        Self::new_with_scope(NetworkScope::V4)
46    }
47
48    /// Creates a new OneShotMdnsResolver with the specified scope
49    pub fn new_with_scope(network_scope: NetworkScope) -> Result<Self, SimpleMdnsError> {
50        Ok(Self {
51            query_timeout: Duration::from_secs(3),
52            unicast_response: UNICAST_RESPONSE,
53            sender_socket: sender_socket(network_scope.is_v4())?,
54            network_scope,
55            receiver_socket: join_multicast(network_scope)?,
56        })
57    }
58
59    /// Send a query packet and returns the first response
60    pub fn query_packet(&self, packet: Packet) -> Result<Option<Vec<u8>>, SimpleMdnsError> {
61        self.sender_socket.send_to(
62            &packet.build_bytes_vec_compressed()?,
63            self.network_scope.socket_address(),
64        )?;
65        let deadline = Instant::now() + self.query_timeout;
66        self.get_next_response(packet.id(), deadline)
67    }
68
69    /// Send a query for A or AAAA (IP v4 and v6 respectively) resources and return the first address
70    pub fn query_service_address(
71        &self,
72        service_name: &str,
73    ) -> Result<Option<std::net::IpAddr>, SimpleMdnsError> {
74        let mut packet = Packet::new_query(0);
75        let service_name = Name::new(service_name)?;
76        packet.questions.push(Question::new(
77            service_name.clone(),
78            TYPE::A.into(),
79            CLASS::IN.into(),
80            self.unicast_response,
81        ));
82
83        self.sender_socket.send_to(
84            &packet.build_bytes_vec_compressed()?,
85            self.network_scope.socket_address(),
86        )?;
87
88        let deadline = Instant::now() + self.query_timeout;
89        loop {
90            let buffer = match self.get_next_response(packet.id(), deadline) {
91                Ok(Some(buffer)) => buffer,
92                Ok(None) => break,
93                Err(err) => {
94                    log::error!("Received invalid packet: {}", err);
95                    continue;
96                }
97            };
98
99            let response = match Packet::parse(&buffer) {
100                Ok(packet) => packet,
101                Err(err) => {
102                    log::error!("Received invalid packet: {}", err);
103                    continue;
104                }
105            };
106
107            for anwser in response.answers {
108                if anwser.name != service_name {
109                    continue;
110                }
111
112                return match anwser.rdata {
113                    RData::A(a) => Ok(Some(IpAddr::V4(Ipv4Addr::from(a.address)))),
114                    RData::AAAA(aaaa) => Ok(Some(IpAddr::V6(Ipv6Addr::from(aaaa.address)))),
115                    _ => Ok(None),
116                };
117            }
118        }
119
120        Ok(None)
121    }
122
123    /// Send a query for SRV resources and return the first address and port
124    pub fn query_service_address_and_port(
125        &self,
126        service_name: &str,
127    ) -> Result<Option<std::net::SocketAddr>, SimpleMdnsError> {
128        let mut packet = Packet::new_query(0);
129        let parsed_name_service = Name::new(service_name)?;
130        packet.questions.push(Question::new(
131            parsed_name_service.clone(),
132            TYPE::SRV.into(),
133            CLASS::IN.into(),
134            self.unicast_response,
135        ));
136
137        self.sender_socket.send_to(
138            &packet.build_bytes_vec()?,
139            self.network_scope.socket_address(),
140        )?;
141
142        let deadline = Instant::now() + self.query_timeout;
143        loop {
144            let buffer = match self.get_next_response(packet.id(), deadline) {
145                Ok(Some(packet)) => packet,
146                Ok(None) => break,
147                Err(err) => {
148                    log::error!("Received invalid packet: {}", err);
149                    continue;
150                }
151            };
152
153            let response = match Packet::parse(&buffer) {
154                Ok(packet) => packet,
155                Err(err) => {
156                    log::error!("Received invalid packet: {}", err);
157                    continue;
158                }
159            };
160
161            let port = response
162                .answers
163                .iter()
164                .filter(|a| a.name == parsed_name_service && a.match_qtype(TYPE::SRV.into()))
165                .find_map(|a| match &a.rdata {
166                    RData::SRV(srv) => Some(srv.port),
167                    _ => None,
168                });
169
170            let mut address = response
171                .additional_records
172                .iter()
173                .filter(|a| a.name == parsed_name_service && a.match_qtype(TYPE::A.into()))
174                .find_map(|a| match &a.rdata {
175                    RData::A(a) => Some(IpAddr::V4(Ipv4Addr::from(a.address))),
176                    RData::AAAA(aaaa) => Some(IpAddr::V6(Ipv6Addr::from(aaaa.address))),
177                    _ => None,
178                });
179
180            if port.is_some() && address.is_none() {
181                address = self.query_service_address(service_name)?;
182            }
183
184            if let (Some(port), Some(address)) = (port, address) {
185                return Ok(Some(SocketAddr::new(address, port)));
186            }
187        }
188
189        Ok(None)
190    }
191
192    /// Set the one shot mdns resolver's query timeout.
193    pub fn set_query_timeout(&mut self, query_timeout: Duration) {
194        self.query_timeout = query_timeout;
195    }
196
197    /// Set the one shot mdns resolver's unicast response.
198    pub fn set_unicast_response(&mut self, unicast_response: bool) {
199        self.unicast_response = unicast_response;
200    }
201
202    fn get_next_response(
203        &self,
204        packet_id: u16,
205        query_deadline: std::time::Instant,
206    ) -> Result<Option<Vec<u8>>, SimpleMdnsError> {
207        let mut buf = [0u8; 4096];
208        loop {
209            match self.receiver_socket.recv_from(&mut buf[..]) {
210                Ok((count, _)) => {
211                    if header_buffer::has_flags(&buf, simple_dns::PacketFlag::RESPONSE)?
212                        && header_buffer::id(&buf)? == packet_id
213                        && header_buffer::answers(&buf)? > 0
214                    {
215                        return Ok(Some(buf[..count].to_vec()));
216                    }
217                }
218                Err(_) => {
219                    if std::time::Instant::now() > query_deadline {
220                        return Ok(None);
221                    }
222                }
223            }
224        }
225    }
226}