Files
lanspread/crates/lanspread-peer/src/services/discovery.rs
T
ddidderr 42cf98ecec fix(peer): charge discovery quota to packet origin
Carry the mDNS response source independently of its advertised A or AAAA target. Active and cooling candidates now charge the observed host, so rotating peer IDs, ports, and target addresses cannot escape the eight-candidate origin budget.

Test Plan:
- just test
- just clippy
- rotating-advertised-target regression
- git diff --check
2026-09-12 13:09:06 +02:00

814 lines
28 KiB
Rust

//! 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<PeerEndpoint, std::net::IpAddr>,
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<PeerId>,
proto_ver: Option<u32>,
}
struct ProtocolNegotiation {
endpoint: PeerEndpoint,
handshake: ReservedCandidateHandshake,
}
struct DiscoveryWorker {
shutdown: CancellationToken,
result_rx: Option<oneshot::Receiver<eyre::Result<()>>>,
thread: Option<JoinHandle<()>>,
}
impl DiscoveryWorker {
fn spawn(
service_type: String,
service_tx: mpsc::Sender<MdnsService>,
shutdown: CancellationToken,
) -> eyre::Result<Self> {
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<Self> {
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<()>>,
) -> 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<PeerEndpoint, std::net::IpAddr>,
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<PeerEndpoint, std::net::IpAddr>,
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<PeerEndpoint, std::net::IpAddr>,
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<MdnsService>,
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::<PeerId>().ok()),
proto_ver: service
.properties
.get("proto_ver")
.and_then(|value| value.parse::<u32>().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<PeerEndpoint> {
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<F>(mut children: FuturesUnordered<F>)
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::<HashMap<_, _>>();
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<AtomicUsize>,
) {
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"
);
}
}