simple_mdns/sync_discovery/
oneshot_resolver.rs1use 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#[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 pub fn new() -> Result<Self, SimpleMdnsError> {
45 Self::new_with_scope(NetworkScope::V4)
46 }
47
48 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 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 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 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 pub fn set_query_timeout(&mut self, query_timeout: Duration) {
194 self.query_timeout = query_timeout;
195 }
196
197 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}