1mod 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 link_local_scopes: std::collections::HashMap<Ipv6Addr, u32>,
49}
50
51const EVENT_CHANNEL_CAPACITY: usize = 256;
52
53pub 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 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 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 dedup_records(&mut all_answers);
149 dedup_records(&mut all_additional);
150 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 let expired = state.cache.evict_expired();
179 for (name, rtype) in expired {
180 let _ = event_tx.send(MdnsEvent::ServiceExpired { name, rtype });
181 }
182
183 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 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 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 let mut mcast_sockets: Vec<McastSocket> = Vec::new();
230
231 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 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 {
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 {
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 pub fn subscribe(&self) -> broadcast::Receiver<MdnsEvent> {
331 self.event_tx.subscribe()
332 }
333
334 pub async fn add_query(&self, label: &str, qtype: u16, interval: Duration) {
336 let mut state = self.inner.lock().await;
337 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 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 pub async fn register_service(&self, reg: ServiceRegistration) {
358 let mut state = self.inner.lock().await;
359 state.services.push(reg);
360 }
361
362 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 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(®, 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 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 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 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 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}