simple_mdns/async_discovery/
oneshot_resolver.rs1use 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
14pub 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 pub fn new() -> Result<Self, SimpleMdnsError> {
49 Self::new_with_scope(NetworkScope::V4)
50 }
51
52 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 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 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 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 pub fn set_query_timeout(&mut self, query_timeout: Duration) {
207 self.query_timeout = query_timeout;
208 }
209
210 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}