//! mDNS peer discovery and discovery-time protocol negotiation. use std::{ collections::{HashMap, VecDeque}, future::Future, thread::JoinHandle, time::Duration, }; use eyre::WrapErr as _; use futures::{StreamExt as _, stream::FuturesUnordered}; use lanspread_mdns::{LANSPREAD_SERVICE_TYPE, MdnsBrowser, MdnsService, MdnsServicePoll}; use lanspread_proto::{PROTOCOL_VERSION, PeerEndpoint, PeerId}; use tokio::sync::{mpsc, oneshot}; use tokio_util::sync::CancellationToken; use crate::{ PeerEvent, PeerEventSender, context::NetworkServiceCtx, events, services::{ handshake::{HandshakeCtx, ReservedCandidateHandshake}, state_sync::run_state_sync, }, }; const MAX_ACTIVE_DISCOVERY_CANDIDATES: usize = 64; /// Upper bound on active-or-cooling candidates observed from one response IP. /// mDNS lets a single LAN host claim arbitrarily many target addresses, peer /// IDs, and ports, so accounting must use packet provenance rather than the /// advertised A/AAAA record. Legitimate hosts run one peer, occasionally a few. const MAX_DISCOVERY_CANDIDATES_PER_SOURCE_IP: usize = 8; const MAX_PENDING_MDNS_SERVICES: usize = 64; const DISCOVERY_CANDIDATE_COOLDOWN: Duration = Duration::from_secs(5); #[derive(Default)] struct RecentCandidates { entries: VecDeque<(PeerEndpoint, std::net::IpAddr, tokio::time::Instant)>, } impl RecentCandidates { fn try_record( &mut self, candidate: PeerEndpoint, source_ip: std::net::IpAddr, now: tokio::time::Instant, ) -> bool { self.expire(now); if self.entries.iter().any(|(endpoint, _, _)| { endpoint.peer_id == candidate.peer_id || endpoint.addr == candidate.addr }) || self.entries.len() >= MAX_ACTIVE_DISCOVERY_CANDIDATES { return false; } self.entries .push_back((candidate, source_ip, now + DISCOVERY_CANDIDATE_COOLDOWN)); true } fn expire(&mut self, now: tokio::time::Instant) { while self .entries .front() .is_some_and(|(_, _, deadline)| *deadline <= now) { self.entries.pop_front(); } } /// Counts unexpired recent candidates from `ip` that are not also in the /// active set, so the caller can sum both without double counting. fn count_from_ip_outside( &mut self, ip: std::net::IpAddr, active: &HashMap, now: tokio::time::Instant, ) -> usize { self.expire(now); self.entries .iter() .filter(|(endpoint, source_ip, _)| *source_ip == ip && !active.contains_key(endpoint)) .count() } #[cfg(test)] fn len(&self) -> usize { self.entries.len() } } struct MdnsPeerInfo { addr: std::net::SocketAddr, source_ip: std::net::IpAddr, peer_id: Option, proto_ver: Option, } struct ProtocolNegotiation { endpoint: PeerEndpoint, handshake: ReservedCandidateHandshake, } struct DiscoveryWorker { shutdown: CancellationToken, result_rx: Option>>, thread: Option>, } impl DiscoveryWorker { fn spawn( service_type: String, service_tx: mpsc::Sender, shutdown: CancellationToken, ) -> eyre::Result { Self::spawn_with(shutdown, move |shutdown| { run_mdns_browser(&service_type, &service_tx, &shutdown) }) } fn spawn_with( shutdown: CancellationToken, worker: impl FnOnce(CancellationToken) -> eyre::Result<()> + Send + 'static, ) -> eyre::Result { let (result_tx, result_rx) = oneshot::channel(); let worker_shutdown = shutdown.clone(); let thread = std::thread::Builder::new() .name("lanspread-mdns-browser".to_owned()) .spawn(move || { let result = worker(worker_shutdown); let _ = result_tx.send(result); }) .wrap_err("failed to spawn mDNS discovery worker")?; Ok(Self { shutdown, result_rx: Some(result_rx), thread: Some(thread), }) } async fn wait_result(&mut self) -> eyre::Result<()> { let result_rx = self .result_rx .as_mut() .ok_or_else(|| eyre::eyre!("mDNS discovery result was already consumed"))?; result_rx .await .map_err(|_| eyre::eyre!("mDNS discovery worker stopped without a result"))? } async fn shutdown_and_join( mut self, observed_result: Option>, ) -> eyre::Result<()> { self.shutdown.cancel(); let result = if let Some(result) = observed_result { self.result_rx.take(); result } else { let result_rx = self .result_rx .take() .ok_or_else(|| eyre::eyre!("mDNS discovery result was already consumed"))?; result_rx .await .map_err(|_| eyre::eyre!("mDNS discovery worker stopped without a result"))? }; self.join_thread()?; result } fn join_thread(&mut self) -> eyre::Result<()> { let Some(thread) = self.thread.take() else { return Ok(()); }; thread .join() .map_err(|_| eyre::eyre!("mDNS discovery worker panicked")) } } impl Drop for DiscoveryWorker { fn drop(&mut self) { self.shutdown.cancel(); if let Err(err) = self.join_thread() { log::error!("Failed to join mDNS discovery worker during cleanup: {err}"); } } } /// Runs the peer discovery service using mDNS. #[allow(clippy::too_many_lines)] pub async fn run_peer_discovery( tx_notify_ui: PeerEventSender, ctx: NetworkServiceCtx, ) -> eyre::Result<()> { log::info!("Starting peer discovery task"); if !wait_for_local_peer_addr(&ctx).await { return Ok(()); } let service_type = LANSPREAD_SERVICE_TYPE.to_string(); let (service_tx, mut service_rx) = tokio::sync::mpsc::channel(MAX_PENDING_MDNS_SERVICES); let service_shutdown = ctx.shutdown.child_token(); let mut worker = DiscoveryWorker::spawn(service_type, service_tx, service_shutdown.clone())?; let mut negotiations = FuturesUnordered::new(); let mut active_candidates = HashMap::new(); let mut recent_candidates = RecentCandidates::default(); let mut mismatch_emitted = false; let mut state_sync = Box::pin(run_state_sync( ctx.clone(), tx_notify_ui.clone(), service_shutdown.clone(), )); let mut observed_state_sync_result = None; let observed_worker_result = loop { tokio::select! { () = ctx.shutdown.cancelled() => break None, result = worker.wait_result() => break Some(result), result = &mut state_sync => { observed_state_sync_result = Some(result); break None; } completed = negotiations.next(), if !negotiations.is_empty() => { if let Some(endpoint) = completed { active_candidates.remove(&endpoint); } } service = service_rx.recv() => { let Some(service) = service else { break None; }; let info = parse_mdns_peer(&service); if is_self_advertisement(&info, &ctx).await { log::trace!("Ignoring self advertisement at {}", info.addr); continue; } if info.proto_ver != Some(PROTOCOL_VERSION) { if !mismatch_emitted { events::send( &tx_notify_ui, PeerEvent::IncompatibleProtocolDetected { observed: info.proto_ver, expected: PROTOCOL_VERSION, }, ); mismatch_emitted = true; } continue; } if let Some(endpoint) = validated_candidate_endpoint(&info) { if !candidate_is_admissible( &active_candidates, &mut recent_candidates, endpoint, info.source_ip, tokio::time::Instant::now(), ) { log::warn!( "Discovery candidate is cooling down or the recent-attempt limit is full; ignoring {}", endpoint.addr ); continue; } let handshake_ctx = HandshakeCtx::from_network(&ctx, &tx_notify_ui) .with_cancellation(service_shutdown.clone()); let handshake = match ReservedCandidateHandshake::reserve(handshake_ctx, endpoint).await { Ok(handshake) => handshake, Err(error) => { log::warn!("Failed to reserve discovery candidate {}: {error}", endpoint.addr); continue; } }; active_candidates.insert(endpoint, info.source_ip); negotiations.push(run_protocol_negotiation(ProtocolNegotiation { endpoint, handshake, })); } } } }; service_shutdown.cancel(); drain_service_children(negotiations).await; let state_sync_exited_early = observed_state_sync_result.is_some(); let state_sync_result = match observed_state_sync_result { Some(result) => result, None => state_sync.await, }; let worker_result = worker.shutdown_and_join(observed_worker_result).await; if let Err(error) = state_sync_result { return Err(error.wrap_err("peer state-sync service failed")); } if state_sync_exited_early && !ctx.shutdown.is_cancelled() { eyre::bail!("peer state-sync service exited unexpectedly"); } match worker_result { Ok(()) if ctx.shutdown.is_cancelled() => Ok(()), Ok(()) => { eyre::bail!("mDNS discovery worker exited unexpectedly"); } Err(err) if ctx.shutdown.is_cancelled() => { log::debug!("Peer discovery worker stopped during shutdown: {err}"); Ok(()) } Err(err) => Err(err.wrap_err("peer discovery worker failed")), } } fn candidate_conflicts( active: &HashMap, candidate: PeerEndpoint, ) -> bool { active .keys() .any(|endpoint| endpoint.peer_id == candidate.peer_id || endpoint.addr == candidate.addr) } /// Returns whether admitting `candidate` would exceed the per-source-IP /// discovery budget across active negotiations and cooling-down attempts. fn source_ip_is_saturated( active: &HashMap, recent: &mut RecentCandidates, source_ip: std::net::IpAddr, now: tokio::time::Instant, ) -> bool { let active_from_ip = active .values() .filter(|observed_ip| **observed_ip == source_ip) .count(); active_from_ip + recent.count_from_ip_outside(source_ip, active, now) >= MAX_DISCOVERY_CANDIDATES_PER_SOURCE_IP } fn candidate_is_admissible( active: &HashMap, recent: &mut RecentCandidates, candidate: PeerEndpoint, source_ip: std::net::IpAddr, now: tokio::time::Instant, ) -> bool { active.len() < MAX_ACTIVE_DISCOVERY_CANDIDATES && !candidate_conflicts(active, candidate) && !source_ip_is_saturated(active, recent, source_ip, now) && recent.try_record(candidate, source_ip, now) } fn run_mdns_browser( service_type: &str, service_tx: &mpsc::Sender, shutdown: &CancellationToken, ) -> eyre::Result<()> { let browser = MdnsBrowser::new(service_type)?; let browse_result = (|| { while !shutdown.is_cancelled() { match browser.next_service_timeout(None, Duration::from_millis(250))? { MdnsServicePoll::Service(service) => { match service_tx.try_send(service) { Ok(()) => {} Err(mpsc::error::TrySendError::Full(_)) => { // Repeated mDNS observations are hints only. Coalesce an // overflow by dropping it rather than letting the native // browser thread allocate without bound or block shutdown. log::trace!( "Coalescing mDNS observation while the bounded discovery queue is full" ); } Err(mpsc::error::TrySendError::Closed(_)) => { log::debug!("Peer discovery consumer dropped; stopping worker"); break; } } } MdnsServicePoll::Timeout => {} MdnsServicePoll::Closed => { log::warn!("mDNS browser closed; stopping peer discovery worker"); break; } } } Ok(()) })(); let close_result = browser.close(); match (browse_result, close_result) { (Ok(()), Ok(())) => Ok(()), (Err(err), Ok(())) | (Ok(()), Err(err)) => Err(err), (Err(browse_err), Err(close_err)) => Err(eyre::eyre!( "mDNS browse failed: {browse_err:#}; browser shutdown also failed: {close_err:#}" )), } } async fn wait_for_local_peer_addr(ctx: &NetworkServiceCtx) -> bool { loop { if ctx.local_peer_addr.read().await.is_some() { return true; } tokio::select! { () = ctx.shutdown.cancelled() => return false, () = tokio::time::sleep(Duration::from_millis(25)) => {} } } } fn parse_mdns_peer(service: &MdnsService) -> MdnsPeerInfo { MdnsPeerInfo { addr: service.addr, source_ip: service.source_ip, peer_id: service .properties .get("peer_id") .and_then(|value| value.parse::().ok()), proto_ver: service .properties .get("proto_ver") .and_then(|value| value.parse::().ok()), } } async fn is_self_advertisement(info: &MdnsPeerInfo, ctx: &NetworkServiceCtx) -> bool { let guard = ctx.local_peer_addr.read().await; guard.as_ref().is_some_and(|addr| *addr == info.addr) || info .peer_id .as_ref() .is_some_and(|peer_id| *peer_id == ctx.peer_id) } /// Returns whether a discovered socket address is a plausible unicast QUIC /// endpoint. mDNS records are attacker-controlled, so multicast, broadcast, /// unspecified, and zero-port targets are dropped before any handshake packet /// is sent toward them. fn is_admissible_candidate_addr(addr: std::net::SocketAddr) -> bool { if addr.port() == 0 { return false; } match addr.ip() { std::net::IpAddr::V4(ip) => { !ip.is_unspecified() && !ip.is_multicast() && !ip.is_broadcast() } std::net::IpAddr::V6(ip) => !ip.is_unspecified() && !ip.is_multicast(), } } fn validated_candidate_endpoint(info: &MdnsPeerInfo) -> Option { if info.proto_ver != Some(PROTOCOL_VERSION) { log::debug!( "Ignoring peer at {} with protocol {:?}; expected {PROTOCOL_VERSION}", info.addr, info.proto_ver ); return None; } if !is_admissible_candidate_addr(info.addr) { log::debug!( "Ignoring current-protocol peer advertised at non-unicast or zero-port address {}", info.addr ); return None; } let Some(peer_id) = info.peer_id else { log::debug!( "Ignoring current-protocol peer at {} without a peer_id TXT record", info.addr ); return None; }; Some(PeerEndpoint::new(peer_id, info.addr)) } async fn run_protocol_negotiation(negotiation: ProtocolNegotiation) -> PeerEndpoint { let endpoint = negotiation.endpoint; let result = negotiation.handshake.run().await; if let Err(err) = result { log::warn!( "Failed to negotiate protocol with peer {}: {err}", endpoint.addr ); } endpoint } async fn drain_service_children(mut children: FuturesUnordered) where F: Future, { while children.next().await.is_some() {} } #[cfg(test)] mod tests { use std::{ collections::HashMap, net::{IpAddr, SocketAddr}, sync::{ Arc, atomic::{AtomicBool, AtomicUsize, Ordering}, mpsc, }, time::Duration, }; use futures::stream::FuturesUnordered; use lanspread_proto::{PeerEndpoint, PeerId}; use tokio_util::sync::CancellationToken; use super::{ DISCOVERY_CANDIDATE_COOLDOWN, DiscoveryWorker, MAX_ACTIVE_DISCOVERY_CANDIDATES, MAX_DISCOVERY_CANDIDATES_PER_SOURCE_IP, MdnsPeerInfo, RecentCandidates, candidate_conflicts, candidate_is_admissible, drain_service_children, is_admissible_candidate_addr, validated_candidate_endpoint, }; fn endpoint(seed: u8, port: u16) -> PeerEndpoint { PeerEndpoint::new( PeerId::from_bytes([seed; 32]), SocketAddr::from(([127, 0, 0, 1], port)), ) } #[test] fn advertised_candidate_addresses_must_be_unicast_with_a_nonzero_port() { for addr in [ "192.168.1.50:42424", "10.0.0.7:1", "127.0.0.1:42424", "[fe80::1]:42424", "[::1]:42424", ] { let addr: SocketAddr = addr.parse().expect("test address parses"); assert!(is_admissible_candidate_addr(addr), "rejected {addr}"); } for addr in [ "192.168.1.50:0", "0.0.0.0:42424", "224.0.0.251:42424", "239.255.255.250:42424", "255.255.255.255:42424", "[::]:42424", "[ff02::fb]:42424", ] { let addr: SocketAddr = addr.parse().expect("test address parses"); assert!(!is_admissible_candidate_addr(addr), "accepted {addr}"); } } #[test] fn validated_candidate_endpoint_drops_non_unicast_advertisements() { let peer_id = Some(PeerId::from_bytes([7; 32])); let unicast = MdnsPeerInfo { addr: "192.168.1.50:42424".parse().expect("test address parses"), source_ip: "192.168.1.50".parse().expect("test source parses"), peer_id, proto_ver: Some(lanspread_proto::PROTOCOL_VERSION), }; assert!(validated_candidate_endpoint(&unicast).is_some()); let multicast = MdnsPeerInfo { addr: "224.0.0.251:42424".parse().expect("test address parses"), source_ip: "192.168.1.50".parse().expect("test source parses"), peer_id, proto_ver: Some(lanspread_proto::PROTOCOL_VERSION), }; assert!(validated_candidate_endpoint(&multicast).is_none()); let zero_port = MdnsPeerInfo { addr: "192.168.1.50:0".parse().expect("test address parses"), source_ip: "192.168.1.50".parse().expect("test source parses"), peer_id, proto_ver: Some(lanspread_proto::PROTOCOL_VERSION), }; assert!(validated_candidate_endpoint(&zero_port).is_none()); } #[test] fn active_candidate_keys_bound_both_claimed_identity_and_address() { let active_endpoint = endpoint(1, 12001); let mut active = HashMap::from([(active_endpoint, active_endpoint.addr.ip())]); assert!(candidate_conflicts(&active, endpoint(1, 12002))); assert!(candidate_conflicts(&active, endpoint(2, 12001))); assert!(!candidate_conflicts(&active, endpoint(2, 12002))); active.remove(&active_endpoint); assert!(!candidate_conflicts(&active, endpoint(1, 12002))); assert!(!candidate_conflicts(&active, endpoint(2, 12001))); } #[test] fn completed_candidate_attempts_are_bounded_and_rate_limited_by_identity_and_address() { let now = tokio::time::Instant::now(); let first = endpoint(1, 12001); let mut recent = RecentCandidates::default(); assert!(recent.try_record(first, first.addr.ip(), now)); for port in 12002..12102 { let candidate = endpoint(1, port); assert!(!recent.try_record(candidate, candidate.addr.ip(), now)); } for seed in 2..=101 { let candidate = endpoint(seed, 12001); assert!(!recent.try_record(candidate, candidate.addr.ip(), now)); } for seed in 2..=u8::try_from(MAX_ACTIVE_DISCOVERY_CANDIDATES).expect("test bound fits u8") { let candidate = endpoint(seed, 13000 + u16::from(seed)); assert!(recent.try_record(candidate, candidate.addr.ip(), now)); } assert_eq!(recent.len(), MAX_ACTIVE_DISCOVERY_CANDIDATES); let last = endpoint(200, 14000); assert!(!recent.try_record(last, last.addr.ip(), now)); let after_cooldown = now + DISCOVERY_CANDIDATE_COOLDOWN; assert!(recent.try_record(last, last.addr.ip(), after_cooldown)); assert_eq!(recent.len(), 1); } fn endpoint_at(seed: u8, ip: [u8; 4], port: u16) -> PeerEndpoint { PeerEndpoint::new(PeerId::from_bytes([seed; 32]), SocketAddr::from((ip, port))) } #[test] fn rotating_advertised_targets_share_the_observed_source_budget() { let now = tokio::time::Instant::now(); let mut active = HashMap::new(); let mut recent = RecentCandidates::default(); let flood_source: IpAddr = [192, 168, 1, 66].into(); // One responder can rotate advertised A records, peer IDs, and ports, // but all candidates remain charged to the observed datagram source. for index in 0..MAX_DISCOVERY_CANDIDATES_PER_SOURCE_IP { let seed = u8::try_from(index + 1).expect("test index fits u8"); let candidate = endpoint_at(seed, [10, 0, seed, 2], 20000 + u16::from(seed)); assert!(candidate_is_admissible( &active, &mut recent, candidate, flood_source, now )); active.insert(candidate, flood_source); } assert!(!candidate_is_admissible( &active, &mut recent, endpoint_at(100, [10, 1, 1, 2], 30000), flood_source, now )); // Other hosts are unaffected while the flood host is saturated. let other_source: IpAddr = [192, 168, 1, 67].into(); assert!(candidate_is_admissible( &active, &mut recent, endpoint_at(101, [10, 1, 1, 3], 42424), other_source, now )); // Completed negotiations keep counting while they cool down ... let completed = endpoint_at(1, [10, 0, 1, 2], 20001); active.remove(&completed); assert!(!candidate_is_admissible( &active, &mut recent, endpoint_at(102, [10, 1, 1, 4], 30001), flood_source, now )); // ... and free the budget once the cooldown expires. let after_cooldown = now + DISCOVERY_CANDIDATE_COOLDOWN; assert!(candidate_is_admissible( &active, &mut recent, endpoint_at(102, [10, 1, 1, 4], 30001), flood_source, after_cooldown )); } #[test] fn expiring_recent_entries_never_bypasses_the_independent_active_cap() { let now = tokio::time::Instant::now(); // Spread the candidates over distinct source hosts so the per-IP // budget does not interfere with the global active cap under test. let mut active = (0..MAX_ACTIVE_DISCOVERY_CANDIDATES) .map(|index| { let seed = u8::try_from(index + 1).expect("test index fits u8"); let candidate = endpoint_at(seed, [10, 0, seed, 1], 15000 + u16::from(seed)); (candidate, candidate.addr.ip()) }) .collect::>(); let mut recent = RecentCandidates::default(); for (candidate, source_ip) in &active { assert!(recent.try_record(*candidate, *source_ip, now)); } let after_cooldown = now + DISCOVERY_CANDIDATE_COOLDOWN; let next = endpoint_at(100, [10, 0, 100, 1], 16000); assert!(!candidate_is_admissible( &active, &mut recent, next, next.addr.ip(), after_cooldown, )); assert_eq!(active.len(), MAX_ACTIVE_DISCOVERY_CANDIDATES); let completed = *active.keys().next().expect("active set should be nonempty"); active.remove(&completed); assert!(candidate_is_admissible( &active, &mut recent, next, next.addr.ip(), after_cooldown, )); } async fn cancellation_aware_child( started: tokio::sync::mpsc::UnboundedSender<()>, shutdown: CancellationToken, completed: Arc, ) { started.send(()).expect("start receiver should remain open"); shutdown.cancelled().await; tokio::task::yield_now().await; completed.fetch_add(1, Ordering::SeqCst); } #[tokio::test] async fn negotiation_batch_drains_started_children_on_shutdown() { let shutdown = CancellationToken::new(); let (started_tx, mut started_rx) = tokio::sync::mpsc::unbounded_channel(); let completed = Arc::new(AtomicUsize::new(0)); let children = FuturesUnordered::new(); for _ in 0..2 { children.push(cancellation_aware_child( started_tx.clone(), shutdown.clone(), completed.clone(), )); } let control_shutdown = shutdown.clone(); let control = async move { for _ in 0..2 { started_rx .recv() .await .expect("every negotiation should start"); } control_shutdown.cancel(); }; tokio::time::timeout(Duration::from_secs(1), async { tokio::join!(drain_service_children(children), control); }) .await .expect("shutdown should drain every negotiation"); assert_eq!(completed.load(Ordering::SeqCst), 2); } #[test] fn dropping_discovery_worker_cancels_and_joins_its_thread() { let shutdown = CancellationToken::new(); let (started_tx, started_rx) = mpsc::sync_channel(0); let stopped = Arc::new(AtomicBool::new(false)); let worker_stopped = stopped.clone(); let worker = DiscoveryWorker::spawn_with(shutdown, move |shutdown| { started_tx .send(()) .expect("test should wait for worker startup"); while !shutdown.is_cancelled() { std::thread::sleep(Duration::from_millis(1)); } worker_stopped.store(true, Ordering::SeqCst); Ok(()) }) .expect("discovery worker should spawn"); started_rx .recv_timeout(Duration::from_secs(1)) .expect("discovery worker should start"); drop(worker); assert!( stopped.load(Ordering::SeqCst), "worker Drop must not return before its thread stops" ); } }