Skip to main content

matc/mdns2/
mod.rs

1//! minimal mDNS service with continuous discovery, record caching, and service registration.
2//!
3//! this provides a long-running service that:
4//! - Runs continuous discovery with periodic re-queries
5//! - Caches discovered records with TTL-based expiration
6//! - Registers local services and responds to incoming mDNS queries
7//! - Emits discovery events via a broadcast channel; call [`MdnsService::subscribe`] to get
8//!   an independent event stream per caller, supporting concurrent discovery operations
9
10mod dnssd;
11mod protocol;
12
13pub use dnssd::{MdnsEvent, ServiceRegistration};
14pub use protocol::{CachedRecord, RecordCache};
15
16use std::collections::HashSet;
17use std::net::{Ipv4Addr, Ipv6Addr};
18use std::sync::Arc;
19use std::time::{Duration, Instant};
20
21use anyhow::Result;
22use tokio::net::UdpSocket;
23use tokio::sync::broadcast;
24use tokio::sync::Mutex;
25use tokio::sync::mpsc::{self, UnboundedSender};
26use tokio_util::sync::CancellationToken;
27
28use crate::mdns;
29use dnssd::{PeriodicQuery, build_service_records, find_matching_services};
30use protocol::{
31    MDNS_ADDR_V4, MDNS_ADDR_V6, McastSocket, SendCommand, build_response,
32    create_multicast_socket_v4, create_multicast_socket_v6, get_local_ips, send_loop,
33};
34
35fn dedup_records(records: &mut Vec<mdns::RR>) {
36    let mut seen = HashSet::new();
37    records.retain(|r| seen.insert(r.clone()));
38}
39
40struct MdnsServiceInner {
41    cache: RecordCache,
42    queries: Vec<PeriodicQuery>,
43    services: Vec<ServiceRegistration>,
44    local_ips_v4: Vec<Ipv4Addr>,
45    local_ips_v6: Vec<Ipv6Addr>,
46    /// Device link-local IPv6 -> index of the interface its mDNS reply arrived
47    /// on. That index is the correct scope_id for sending to `fe80::...`.
48    link_local_scopes: std::collections::HashMap<Ipv6Addr, u32>,
49}
50
51const EVENT_CHANNEL_CAPACITY: usize = 256;
52
53/// Long-running mDNS service with discovery, caching, and service registration.
54pub struct MdnsService {
55    inner: Arc<Mutex<MdnsServiceInner>>,
56    send_tx: UnboundedSender<SendCommand>,
57    event_tx: broadcast::Sender<MdnsEvent>,
58    cancel: CancellationToken,
59}
60
61async fn recv_loop(
62    socket: Arc<UdpSocket>,
63    interface: Option<u32>,
64    inner: Arc<Mutex<MdnsServiceInner>>,
65    send_tx: UnboundedSender<SendCommand>,
66    event_tx: broadcast::Sender<MdnsEvent>,
67    cancel: CancellationToken,
68) {
69    let mut buf = vec![0u8; 9000];
70    loop {
71        let (n, addr) = tokio::select! {
72            result = socket.recv_from(&mut buf) => {
73                match result {
74                    Ok(v) => v,
75                    Err(e) => {
76                        log::debug!("mdns2 recv error: {}", e);
77                        continue;
78                    }
79                }
80            }
81            _ = cancel.cancelled() => return,
82        };
83
84        let data = &buf[..n];
85        let msg = match mdns::parse_dns(data, addr) {
86            Ok(m) => m,
87            Err(e) => {
88                log::trace!("mdns2: failed to parse packet from {}: {:?}", addr, e);
89                continue;
90            }
91        };
92
93        let is_response = msg.flags & 0x8000 != 0;
94
95        if is_response {
96            // Ingest all records into cache
97            let mut state = inner.lock().await;
98            let all_records: Vec<mdns::RR> = msg
99                .answers
100                .iter()
101                .chain(msg.additional.iter())
102                .cloned()
103                .collect();
104
105            let mut new_ptr_records = Vec::new();
106            for rr in &all_records {
107                state.cache.ingest(rr);
108                // Remember the receive interface for link-local AAAA records
109                if let (Some(idx), mdns::RRData::AAAA(ip)) = (interface, &rr.data) {
110                    if (ip.segments()[0] & 0xffc0) == 0xfe80 {
111                        state.link_local_scopes.insert(*ip, idx);
112                    }
113                }
114                if rr.typ == mdns::TYPE_PTR {
115                    if let mdns::RRData::PTR(ref target) = rr.data {
116                        new_ptr_records.push((rr.name.clone(), target.clone()));
117                    }
118                }
119            }
120            for (name, target) in new_ptr_records {
121                let _ = event_tx.send(MdnsEvent::ServiceDiscovered {
122                    name,
123                    target,
124                    records: all_records.clone(),
125                });
126            }
127        } else {
128            let state = inner.lock().await;
129            if state.services.is_empty() {
130                continue;
131            }
132            let mut all_answers = Vec::new();
133            let mut all_additional = Vec::new();
134            for q in &msg.queries {
135                let (ans, add) = find_matching_services(
136                    &q.name,
137                    q.typ,
138                    &state.services,
139                    &state.local_ips_v4,
140                    &state.local_ips_v6,
141                );
142                all_answers.extend(ans);
143                all_additional.extend(add);
144            }
145            drop(state);
146
147            // Deduplicate records that matched multiple queries
148            dedup_records(&mut all_answers);
149            dedup_records(&mut all_additional);
150            // Don't repeat answer records in additional
151            all_additional.retain(|r| !all_answers.contains(r));
152
153            if !all_answers.is_empty() {
154                if let Ok(packet) = build_response(&all_answers, &all_additional) {
155                    let _ = send_tx.send(SendCommand::Multicast(packet));
156                }
157            }
158        }
159    }
160}
161
162async fn periodic_loop(
163    inner: Arc<Mutex<MdnsServiceInner>>,
164    send_tx: UnboundedSender<SendCommand>,
165    event_tx: broadcast::Sender<MdnsEvent>,
166    cancel: CancellationToken,
167) {
168    let mut interval = tokio::time::interval(Duration::from_secs(1));
169    loop {
170        tokio::select! {
171            _ = interval.tick() => {}
172            _ = cancel.cancelled() => return,
173        }
174
175        let mut state = inner.lock().await;
176
177        // Evict expired cache entries
178        let expired = state.cache.evict_expired();
179        for (name, rtype) in expired {
180            let _ = event_tx.send(MdnsEvent::ServiceExpired { name, rtype });
181        }
182
183        // Send due queries
184        let now = Instant::now();
185        let mut packets = Vec::new();
186        for q in &mut state.queries {
187            if now.duration_since(q.last_sent) >= q.interval {
188                if let Ok(pkt) = mdns::create_query(&q.label, q.qtype) {
189                    packets.push(pkt);
190                }
191                q.last_sent = now;
192            }
193        }
194        drop(state);
195
196        for pkt in packets {
197            let _ = send_tx.send(SendCommand::Multicast(pkt));
198        }
199
200        // Refresh local IPs periodically (cheap operation)
201        let (v4, v6) = get_local_ips();
202        let mut state = inner.lock().await;
203        state.local_ips_v4 = v4;
204        state.local_ips_v6 = v6;
205    }
206}
207
208impl MdnsService {
209    /// Create a new mDNS service.
210    ///
211    /// Call [`subscribe`](Self::subscribe) on the returned handle to receive discovery events.
212    /// Multiple independent subscribers may receive events concurrently.
213    pub async fn new() -> Result<Arc<Self>> {
214        let (event_tx, _) = broadcast::channel(EVENT_CHANNEL_CAPACITY);
215        let (send_tx, send_rx) = mpsc::unbounded_channel();
216        let cancel = CancellationToken::new();
217
218        let (v4, v6) = get_local_ips();
219        let inner = Arc::new(Mutex::new(MdnsServiceInner {
220            cache: RecordCache::new(),
221            queries: Vec::new(),
222            services: Vec::new(),
223            local_ips_v4: v4,
224            local_ips_v6: v6,
225            link_local_scopes: std::collections::HashMap::new(),
226        }));
227
228        // Create sockets
229        let mut mcast_sockets: Vec<McastSocket> = Vec::new();
230
231        // IPv4
232        match create_multicast_socket_v4() {
233            Ok(std_sock) => match UdpSocket::from_std(std_sock) {
234                Ok(s) => mcast_sockets.push(McastSocket {
235                    sock: Arc::new(s),
236                    multicast_addr: MDNS_ADDR_V4,
237                    interface: None,
238                }),
239                Err(e) => log::warn!("mdns2: failed to wrap v4 socket: {}", e),
240            },
241            Err(e) => log::warn!("mdns2: failed to create v4 socket: {}", e),
242        }
243
244        if let Ok(ifaces) = if_addrs::get_if_addrs() {
245            let mut seen_indices = std::collections::HashSet::new();
246            for iface in ifaces {
247                if !iface.ip().is_ipv6() {
248                    continue;
249                }
250                if let Some(idx) = iface.index {
251                    if !seen_indices.insert(idx) {
252                        continue;
253                    }
254                    match create_multicast_socket_v6(idx) {
255                        Ok(std_sock) => match UdpSocket::from_std(std_sock) {
256                            Ok(s) => mcast_sockets.push(McastSocket {
257                                sock: Arc::new(s),
258                                multicast_addr: MDNS_ADDR_V6,
259                                interface: Some(idx),
260                            }),
261                            Err(e) => {
262                                log::debug!("mdns2: failed to wrap v6 socket idx={}: {}", idx, e)
263                            }
264                        },
265                        Err(e) => {
266                            log::debug!("mdns2: failed to create v6 socket idx={}: {}", idx, e)
267                        }
268                    }
269                }
270            }
271        }
272
273        if mcast_sockets.is_empty() {
274            anyhow::bail!("mdns2: no sockets could be created");
275        }
276
277        // Spawn recv loops (one per socket)
278        for ms in &mcast_sockets {
279            let sock = ms.sock.clone();
280            let interface = ms.interface;
281            let inner = inner.clone();
282            let send_tx = send_tx.clone();
283            let event_tx = event_tx.clone();
284            let cancel = cancel.child_token();
285            tokio::spawn(async move {
286                recv_loop(sock, interface, inner, send_tx, event_tx, cancel).await;
287            });
288        }
289
290        // Spawn periodic loop
291        {
292            let inner = inner.clone();
293            let send_tx = send_tx.clone();
294            let event_tx = event_tx.clone();
295            let cancel = cancel.child_token();
296            tokio::spawn(async move {
297                periodic_loop(inner, send_tx, event_tx, cancel).await;
298            });
299        }
300
301        // Spawn send loop
302        {
303            let cancel = cancel.child_token();
304            tokio::spawn(async move {
305                send_loop(mcast_sockets, send_rx, cancel).await;
306            });
307        }
308
309        let service = Arc::new(MdnsService {
310            inner,
311            send_tx,
312            event_tx,
313            cancel,
314        });
315
316        Ok(service)
317    }
318
319    pub async fn scope_for(&self, ip: &Ipv6Addr) -> Option<u32> {
320        self.inner.lock().await.link_local_scopes.get(ip).copied()
321    }
322
323    /// Subscribe to discovery events.
324    ///
325    /// Returns an independent [`broadcast::Receiver`]; each subscriber receives every event.
326    /// Subscribe before calling [`active_lookup`](Self::active_lookup) to avoid missing
327    /// responses that arrive before the next `recv()` call.
328    /// On lag (`RecvError::Lagged`), log a warning and keep draining — events are recoverable
329    /// by re-issuing [`active_lookup`](Self::active_lookup).
330    pub fn subscribe(&self) -> broadcast::Receiver<MdnsEvent> {
331        self.event_tx.subscribe()
332    }
333
334    /// Add a periodic query. The query will be sent immediately, then every interval.
335    pub async fn add_query(&self, label: &str, qtype: u16, interval: Duration) {
336        let mut state = self.inner.lock().await;
337        // Send immediately
338        let sent_at = Instant::now();
339        if let Ok(pkt) = mdns::create_query(label, qtype) {
340            let _ = self.send_tx.send(SendCommand::Multicast(pkt));
341        }
342        state.queries.push(PeriodicQuery {
343            label: label.to_owned(),
344            qtype,
345            interval,
346            last_sent: sent_at,
347        });
348    }
349
350    /// Remove a periodic query by label.
351    pub async fn remove_query(&self, label: &str) {
352        let mut state = self.inner.lock().await;
353        state.queries.retain(|q| q.label != label);
354    }
355
356    /// Register a local service to be advertised.
357    pub async fn register_service(&self, reg: ServiceRegistration) {
358        let mut state = self.inner.lock().await;
359        state.services.push(reg);
360    }
361
362    /// Unregister a local service. Sends a goodbye (TTL=0) for the service records.
363    pub async fn unregister_service(&self, instance: &str, service_type: &str) {
364        let mut state = self.inner.lock().await;
365        let idx = state
366            .services
367            .iter()
368            .position(|s| s.instance_name == instance && s.service_type == service_type);
369        if let Some(idx) = idx {
370            let reg = state.services.remove(idx);
371            // Build goodbye records (TTL=0)
372            let svc_v4 = reg.ips_v4.as_deref().unwrap_or(&state.local_ips_v4);
373            let svc_v6 = reg.ips_v6.as_deref().unwrap_or(&state.local_ips_v6);
374            let mut goodbye_records = build_service_records(&reg, svc_v4, svc_v6);
375            for rr in &mut goodbye_records {
376                rr.ttl = 0;
377            }
378            drop(state);
379            if let Ok(pkt) = build_response(&goodbye_records, &[]) {
380                let _ = self.send_tx.send(SendCommand::Multicast(pkt));
381            }
382        }
383    }
384
385    /// Send a gratuitous announcement of all registered services.
386    pub async fn announce(&self) {
387        let state = self.inner.lock().await;
388        let mut all_answers = Vec::new();
389        let mut all_additional = Vec::new();
390        for reg in &state.services {
391            let svc_v4 = reg.ips_v4.as_deref().unwrap_or(&state.local_ips_v4);
392            let svc_v6 = reg.ips_v6.as_deref().unwrap_or(&state.local_ips_v6);
393            let records = build_service_records(reg, svc_v4, svc_v6);
394            // PTR goes as answer, everything else as additional
395            for r in records {
396                if r.typ == mdns::TYPE_PTR {
397                    all_answers.push(r);
398                } else {
399                    all_additional.push(r);
400                }
401            }
402        }
403        drop(state);
404
405        if !all_answers.is_empty() {
406            if let Ok(pkt) = build_response(&all_answers, &all_additional) {
407                let _ = self.send_tx.send(SendCommand::Multicast(pkt));
408            }
409        }
410    }
411
412    /// Lookup cached records by name and type.
413    pub async fn lookup(&self, name: &str, qtype: u16) -> Vec<mdns::RR> {
414        let state = self.inner.lock().await;
415        if qtype == mdns::QTYPE_ANY {
416            state.cache.lookup_name(name)
417        } else {
418            state.cache.lookup(name, qtype)
419        }
420    }
421
422    pub async fn active_lookup(&self, name: &str, qtype: u16) {
423        if let Ok(pkt) = mdns::create_query(name, qtype) {
424            let _ = self.send_tx.send(SendCommand::Multicast(pkt));
425        }
426    }
427
428    /// Shut down all background tasks.
429    pub fn shutdown(&self) {
430        self.cancel.cancel();
431    }
432}
433
434impl Drop for MdnsService {
435    fn drop(&mut self) {
436        self.cancel.cancel();
437    }
438}