simple_mdns/sync_discovery/
service_discovery.rs1use simple_dns::{rdata::RData, Name, Packet, Question, CLASS, TYPE};
2
3use std::{
4 collections::HashSet,
5 error::Error,
6 net::{SocketAddr, UdpSocket},
7 sync::{Arc, RwLock},
8 time::{Duration, Instant},
9};
10
11use crate::{
12 resource_record_manager::{
13 service_discovery_resource_manager, DomainResourceFilter, ResourceRecordManager,
14 },
15 InstanceInformation, NetworkScope, SimpleMdnsError,
16};
17
18pub struct ServiceDiscovery {
38 instance_name: Name<'static>,
39 service_name: Name<'static>,
40 resource_manager: Arc<RwLock<ResourceRecordManager<'static>>>,
41 sender_socket: UdpSocket,
42 network_scope: NetworkScope,
43}
44
45impl ServiceDiscovery {
46 pub fn new(
51 instance_information: InstanceInformation,
52 service_name: &str,
53 resource_ttl: u32,
54 ) -> Result<Self, SimpleMdnsError> {
55 Self::new_with_scope(
56 instance_information,
57 service_name,
58 resource_ttl,
59 None,
60 NetworkScope::V4,
61 )
62 }
63
64 pub fn new_with_scope(
72 instance_information: InstanceInformation,
73 service_name: &str,
74 resource_ttl: u32,
75 on_discovery: Option<std::sync::mpsc::Sender<InstanceInformation>>,
76 network_scope: NetworkScope,
77 ) -> Result<Self, SimpleMdnsError> {
78 let instance_full_name = format!(
79 "{}.{service_name}",
80 instance_information.escaped_instance_name()
81 );
82 let instance_full_name = Name::new(&instance_full_name)?.into_owned();
83 let service_name = Name::new(service_name)?.into_owned();
84
85 let resource_manager = service_discovery_resource_manager(
86 &service_name,
87 &instance_full_name,
88 resource_ttl,
89 instance_information,
90 )?;
91
92 let service_discovery = Self {
93 instance_name: instance_full_name,
94 service_name,
95 resource_manager: Arc::new(RwLock::new(resource_manager)),
96 sender_socket: crate::socket_helper::sender_socket(network_scope.is_v4())?,
97 network_scope,
98 };
99
100 service_discovery.receive_packets_loop(on_discovery)?;
101 service_discovery.refresh_known_instances()?;
102 service_discovery.announce(false);
103
104 if let Err(err) = query_service_instances(
105 service_discovery.service_name.clone(),
106 &service_discovery.sender_socket,
107 &service_discovery.network_scope.socket_address(),
108 ) {
109 log::error!("There was an error queruing service instances: {err}");
110 }
111
112 Ok(service_discovery)
113 }
114
115 pub fn remove_service_from_discovery(&mut self) {
118 self.announce(true);
119 self.resource_manager
120 .write()
121 .unwrap()
122 .remove_domain_resources(&self.instance_name);
123 }
124
125 pub fn get_known_services(&self) -> HashSet<InstanceInformation> {
127 self.resource_manager
128 .read()
129 .unwrap()
130 .get_domain_resources(&self.service_name, DomainResourceFilter::cached())
131 .filter_map(|domain_resources| {
132 InstanceInformation::from_records(&self.service_name, domain_resources)
133 })
134 .collect()
135 }
136
137 fn refresh_known_instances(&self) -> std::io::Result<()> {
138 let service_name = self.service_name.clone();
139 let resource_manager = self.resource_manager.clone();
140
141 let sender = self.sender_socket.try_clone()?;
142 let address = self.network_scope.socket_address();
143
144 std::thread::spawn(move || loop {
145 log::info!("Refreshing known services");
146 let now = Instant::now();
147 let next_expiration = resource_manager.read().unwrap().get_next_refresh();
148
149 log::trace!("next expiration: {:?}", next_expiration);
150 match next_expiration {
151 Some(expiration) => {
152 if expiration <= now {
153 if let Err(err) =
154 query_service_instances(service_name.clone(), &sender, &address)
155 {
156 log::error!("There was an error querying service instances. {err}");
157 }
158 std::thread::sleep(Duration::from_secs(5));
159 } else {
160 std::thread::sleep(expiration - now);
161 }
162 }
163 None => {
164 std::thread::sleep(Duration::from_secs(5));
165 }
166 }
167 });
168
169 Ok(())
170 }
171
172 pub fn announce(&self, cache_flush: bool) {
179 let mut packet = Packet::new_reply(1);
180 let resource_manager = self.resource_manager.read().unwrap();
181 let mut additional_records = HashSet::new();
182
183 for d_resources in resource_manager.get_domain_resources(
184 &self.instance_name.clone(),
185 DomainResourceFilter::authoritative(true),
186 ) {
187 if cache_flush {
188 d_resources
189 .filter(|r| r.match_qclass(CLASS::IN.into()))
190 .for_each(|r| packet.answers.push(r.to_cache_flush_record()));
191 } else {
192 d_resources.cloned().for_each(|resource| {
193 if let RData::SRV(srv) = &resource.rdata {
194 let target = resource_manager
195 .get_domain_resources(
196 &srv.target,
197 DomainResourceFilter::authoritative(false),
198 )
199 .flatten()
200 .filter(|r| {
201 (r.match_qtype(TYPE::A.into()) || r.match_qtype(TYPE::AAAA.into()))
202 && r.match_qclass(CLASS::IN.into())
203 })
204 .cloned();
205
206 additional_records.extend(target);
207 }
208
209 packet.answers.push(resource);
210 });
211 };
212 }
213
214 for additional_record in additional_records {
215 packet.additional_records.push(additional_record)
216 }
217
218 if !packet.answers.is_empty()
219 && packet
220 .build_bytes_vec_compressed()
221 .map(|bytes| {
222 send_packet(
223 &self.sender_socket,
224 &bytes,
225 &self.network_scope.socket_address(),
226 )
227 })
228 .is_err()
229 {
230 log::info!("Failed to advertise service");
231 }
232 }
233
234 fn receive_packets_loop(
235 &self,
236 mut on_discovery: Option<std::sync::mpsc::Sender<InstanceInformation>>,
237 ) -> Result<(), SimpleMdnsError> {
238 let service_name = self.service_name.clone();
239 let full_name = self.instance_name.clone();
240 let resources = self.resource_manager.clone();
241 let multicast_address = self.network_scope.socket_address();
242
243 let sender_socket = self.sender_socket.try_clone()?;
244 let recv_socket = crate::socket_helper::join_multicast(self.network_scope)?;
245 recv_socket.set_read_timeout(None)?;
246
247 std::thread::spawn(move || loop {
248 let mut recv_buffer = [0u8; 9000];
249 let (count, addr) = match recv_socket.recv_from(&mut recv_buffer) {
250 Ok(received) => received,
251 Err(err) => {
252 log::error!("Failed to read network information {err}");
253 continue;
254 }
255 };
256
257 match Packet::parse(&recv_buffer[..count]) {
258 Ok(packet) => {
259 if packet.has_flags(simple_dns::PacketFlag::RESPONSE) {
260 add_response_to_resources(
261 packet,
262 &service_name,
263 &full_name,
264 &mut resources.write().unwrap(),
265 &mut on_discovery,
266 )
267 } else {
268 match crate::build_reply(packet, &resources.read().unwrap()) {
269 Some((reply_packet, unicast_response)) => {
270 let reply = match reply_packet.build_bytes_vec_compressed() {
271 Ok(reply) => reply,
272 Err(err) => {
273 log::error!("Failed to build reply {err}");
274 continue;
275 }
276 };
277
278 let reply_addr = if unicast_response {
279 addr
280 } else {
281 multicast_address
282 };
283
284 log::debug!("sending reply");
285 send_packet(&sender_socket, &reply, &reply_addr);
286 }
287 None => {
288 log::debug!("No reply to send");
289 }
290 }
291 }
292 }
293 Err(err) => {
294 log::error!("Received Invalid Packet {err}");
295 }
296 }
297 });
298
299 Ok(())
300 }
301}
302
303fn query_service_instances(
304 service_name: Name,
305 socket: &UdpSocket,
306 address: &SocketAddr,
307) -> Result<(), Box<dyn Error>> {
308 log::trace!("probing service instances");
309 let mut packet = Packet::new_query(0);
310 packet.questions.push(Question::new(
312 service_name,
313 TYPE::PTR.into(),
314 CLASS::IN.into(),
315 false,
316 ));
317
318 send_packet(socket, &packet.build_bytes_vec_compressed()?, address);
319
320 Ok(())
321}
322
323fn send_packet(socket: &UdpSocket, packet_bytes: &[u8], address: &SocketAddr) {
324 if let Err(err) = socket.send_to(packet_bytes, address) {
325 log::error!("There was an error sending the packet: {err}");
326 }
327}
328
329fn add_response_to_resources(
330 packet: Packet,
331 service_name: &Name<'_>,
332 full_name: &Name<'_>,
333 owned_resources: &mut ResourceRecordManager,
334 on_discovery: &mut Option<std::sync::mpsc::Sender<InstanceInformation>>,
335) {
336 let resources = packet
337 .answers
338 .into_iter()
339 .chain(packet.additional_records)
340 .filter(|aw| aw.name.ne(full_name) && aw.name.is_subdomain_of(service_name))
341 .map(|r| r.into_owned());
342
343 if let Some(channel) = on_discovery {
344 let resources: Vec<_> = resources.collect();
345 if resources.is_empty() {
346 return;
347 }
348
349 if let Some(instance_information) =
350 InstanceInformation::from_records(service_name, resources.iter())
351 {
352 if channel.send(instance_information).is_err() {
353 *on_discovery = None
354 }
355 }
356
357 for resource in resources {
358 owned_resources.add_cached_resource(resource);
359 }
360 } else {
361 for resource in resources {
362 owned_resources.add_cached_resource(resource);
363 }
364 }
365}