Skip to main content

simple_mdns/sync_discovery/
service_discovery.rs

1use 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
18/// Service Discovery implementation using DNS-SD.
19/// This implementation advertise all the registered addresses, query for the same service on the same network and
20/// keeps a cache of known service instances
21///
22/// Notice that this crate does not provide any means of finding your own ip address. There are crates that provide this kind of feature.
23///
24/// ## Example
25/// ```
26/// use simple_mdns::sync_discovery::ServiceDiscovery;
27/// use simple_mdns::InstanceInformation;
28/// use std::str::FromStr;
29///
30/// let mut discovery = ServiceDiscovery::new(
31///     InstanceInformation::new("a".into()).with_socket_address("192.168.1.22:8090".parse().expect("Invalid Socket Address")),
32///     "_mysrv._tcp.local",
33///     60
34/// ).expect("Failed to create service discovery");
35///
36/// ```
37pub 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    /// Creates a new ServiceDiscovery by providing `instance_information`, `service_name`, `resource ttl`. The service will be created using IPV4 scope with UNSPECIFIED Interface
47    ///
48    /// `service_name` must be in the standard specified by the mdns RFC, example: **_my_service._tcp.local**
49    /// `resource_ttl` refers to the amount of time in seconds your service will be cached in the dns responder.
50    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    /// Creates a new ServiceDiscovery by providing `instance_information`, `service_name`, `resource ttl`, `on_disovery` and `network_scope`
65    ///
66    /// `service_name` must be in the standard specified by the mdns RFC, example: **_my_service._tcp.local**
67    /// `resource_ttl` refers to the amount of time in seconds your service will be cached in the dns responder.
68    /// `on_discovery` channel, if provided, will receive every instance information when
69    /// discovered
70    /// `network_scope` to be used
71    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    /// Remove service from discovery by announcing with a cache flush and
116    /// removing all the internal resource records
117    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    /// Return the [`InstanceInformation`] of all known services
126    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    /// Announce the service by sending a packet with all the resource records in the answers
173    /// section. It is not necessary to call this method manually, it will be called automatically
174    /// when the instance is added to the discovery.
175    ///
176    /// if `cache_flush` is true, then the resources will have the cache flush flag set, this will
177    /// cause them to be removed from any cache that receives the packet.
178    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    // RFC 6763 §4 — discover service instances via PTR query on the service type name
311    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}