Skip to main content

simple_mdns/async_discovery/
oneshot_resolver.rs

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