use std::{
collections::{BTreeSet, HashMap},
net::{IpAddr, SocketAddr},
str::FromStr,
};
use iroh_base::{NodeId, PublicKey};
use n0_future::{
boxed::BoxStream,
task::{self, AbortOnDropHandle, JoinSet},
time::{self, Duration},
};
use n0_watcher::{Watchable, Watcher as _};
use swarm_discovery::{Discoverer, DropGuard, IpClass, Peer};
use tokio::sync::mpsc::{self, error::TrySendError};
use tracing::{Instrument, debug, error, info_span, trace, warn};
use super::{DiscoveryContext, DiscoveryError, IntoDiscovery, IntoDiscoveryError};
use crate::discovery::{Discovery, DiscoveryEvent, DiscoveryItem, NodeData, NodeInfo};
const N0_LOCAL_SWARM: &str = "iroh.local.swarm";
pub const NAME: &str = "local.swarm.discovery";
const USER_DATA_ATTRIBUTE: &str = "user-data";
const DISCOVERY_DURATION: Duration = Duration::from_secs(10);
#[derive(Debug)]
pub struct MdnsDiscovery {
#[allow(dead_code)]
handle: AbortOnDropHandle<()>,
sender: mpsc::Sender<Message>,
advertise: bool,
local_addrs: Watchable<Option<NodeData>>,
}
#[derive(Debug)]
enum Message {
Discovery(String, Peer),
Resolve(NodeId, mpsc::Sender<Result<DiscoveryItem, DiscoveryError>>),
Timeout(NodeId, usize),
Subscribe(mpsc::Sender<DiscoveryEvent>),
}
#[derive(Debug)]
struct Subscribers(Vec<mpsc::Sender<DiscoveryEvent>>);
impl Subscribers {
fn new() -> Self {
Self(vec![])
}
fn push(&mut self, subscriber: mpsc::Sender<DiscoveryEvent>) {
self.0.push(subscriber);
}
fn send(&mut self, item: DiscoveryEvent) {
let mut clean_up = vec![];
for (i, subscriber) in self.0.iter().enumerate() {
if let Err(err) = subscriber.try_send(item.clone()) {
match err {
TrySendError::Full(_) => {
warn!(
?item,
idx = i,
"local swarm discovery subscriber is blocked, dropping item"
)
}
TrySendError::Closed(_) => clean_up.push(i),
}
}
}
for i in clean_up.into_iter().rev() {
self.0.swap_remove(i);
}
}
}
#[derive(Debug)]
pub struct MdnsDiscoveryBuilder {
advertise: bool,
}
impl MdnsDiscoveryBuilder {
pub fn new() -> Self {
Self { advertise: true }
}
pub fn advertise(mut self, advertise: bool) -> Self {
self.advertise = advertise;
self
}
pub fn build(self, node_id: NodeId) -> Result<MdnsDiscovery, IntoDiscoveryError> {
MdnsDiscovery::new(node_id, self.advertise)
}
}
impl Default for MdnsDiscoveryBuilder {
fn default() -> Self {
Self::new()
}
}
impl IntoDiscovery for MdnsDiscoveryBuilder {
fn into_discovery(
self,
context: &DiscoveryContext,
) -> Result<impl Discovery, IntoDiscoveryError> {
self.build(context.node_id())
}
}
impl MdnsDiscovery {
pub fn builder() -> MdnsDiscoveryBuilder {
MdnsDiscoveryBuilder::new()
}
pub fn new(node_id: NodeId, advertise: bool) -> Result<Self, IntoDiscoveryError> {
debug!("Creating new MdnsDiscovery service");
let (send, mut recv) = mpsc::channel(64);
let task_sender = send.clone();
let rt = tokio::runtime::Handle::current();
let discovery = MdnsDiscovery::spawn_discoverer(
node_id,
advertise,
task_sender.clone(),
BTreeSet::new(),
&rt,
)?;
let local_addrs: Watchable<Option<NodeData>> = Watchable::default();
let mut addrs_change = local_addrs.watch();
let discovery_fut = async move {
let mut node_addrs: HashMap<PublicKey, Peer> = HashMap::default();
let mut subscribers = Subscribers::new();
let mut last_id = 0;
let mut senders: HashMap<
PublicKey,
HashMap<usize, mpsc::Sender<Result<DiscoveryItem, DiscoveryError>>>,
> = HashMap::default();
let mut timeouts = JoinSet::new();
loop {
trace!(?node_addrs, "MdnsDiscovery Service loop tick");
let msg = tokio::select! {
msg = recv.recv() => {
msg
}
Ok(Some(data)) = addrs_change.updated() => {
tracing::trace!(?data, "MdnsDiscovery address changed");
discovery.remove_all();
let addrs =
MdnsDiscovery::socketaddrs_to_addrs(data.direct_addresses());
for addr in addrs {
discovery.add(addr.0, addr.1)
}
if let Some(user_data) = data.user_data() {
if let Err(err) = discovery.set_txt_attribute(USER_DATA_ATTRIBUTE.to_string(), Some(user_data.to_string())) {
warn!("Failed to set the user-defined data in local swarm discovery: {err:?}");
}
}
continue;
}
};
let msg = match msg {
None => {
error!("MdnsDiscovery channel closed");
error!("closing MdnsDiscovery");
timeouts.abort_all();
discovery.remove_all();
return;
}
Some(msg) => msg,
};
match msg {
Message::Discovery(discovered_node_id, peer_info) => {
trace!(
?discovered_node_id,
?peer_info,
"MdnsDiscovery Message::Discovery"
);
let discovered_node_id = match PublicKey::from_str(&discovered_node_id) {
Ok(node_id) => node_id,
Err(e) => {
warn!(
discovered_node_id,
"couldn't parse node_id from mdns discovery service: {e:?}"
);
continue;
}
};
if discovered_node_id == node_id {
continue;
}
if peer_info.is_expiry() {
trace!(
?discovered_node_id,
"removing node from MdnsDiscovery address book"
);
node_addrs.remove(&discovered_node_id);
subscribers.send(DiscoveryEvent::Expired(discovered_node_id));
continue;
}
let entry = node_addrs.entry(discovered_node_id);
if let std::collections::hash_map::Entry::Occupied(ref entry) = entry {
if entry.get() == &peer_info {
continue;
}
}
debug!(
?discovered_node_id,
?peer_info,
"adding node to MdnsDiscovery address book"
);
let mut resolved = false;
let item = peer_to_discovery_item(&peer_info, &discovered_node_id);
if let Some(senders) = senders.get(&discovered_node_id) {
trace!(?item, senders = senders.len(), "sending DiscoveryItem");
resolved = true;
for sender in senders.values() {
sender.send(Ok(item.clone())).await.ok();
}
}
entry.or_insert(peer_info);
if !resolved {
subscribers.send(DiscoveryEvent::Discovered(item));
}
}
Message::Resolve(node_id, sender) => {
let id = last_id + 1;
last_id = id;
trace!(?node_id, "MdnsDiscovery Message::SendAddrs");
if let Some(peer_info) = node_addrs.get(&node_id) {
let item = peer_to_discovery_item(peer_info, &node_id);
debug!(?item, "sending DiscoveryItem");
sender.send(Ok(item)).await.ok();
}
if let Some(senders_for_node_id) = senders.get_mut(&node_id) {
senders_for_node_id.insert(id, sender);
} else {
let mut senders_for_node_id = HashMap::new();
senders_for_node_id.insert(id, sender);
senders.insert(node_id, senders_for_node_id);
}
let timeout_sender = task_sender.clone();
timeouts.spawn(async move {
time::sleep(DISCOVERY_DURATION).await;
trace!(?node_id, "discovery timeout");
timeout_sender
.send(Message::Timeout(node_id, id))
.await
.ok();
});
}
Message::Timeout(node_id, id) => {
trace!(?node_id, "MdnsDiscovery Message::Timeout");
if let Some(senders_for_node_id) = senders.get_mut(&node_id) {
senders_for_node_id.remove(&id);
if senders_for_node_id.is_empty() {
senders.remove(&node_id);
}
}
}
Message::Subscribe(subscriber) => {
trace!("MdnsDiscovery Message::Subscribe");
subscribers.push(subscriber);
}
}
}
};
let handle = task::spawn(discovery_fut.instrument(info_span!("swarm-discovery.actor")));
Ok(Self {
handle: AbortOnDropHandle::new(handle),
sender: send,
advertise,
local_addrs,
})
}
fn spawn_discoverer(
node_id: PublicKey,
advertise: bool,
sender: mpsc::Sender<Message>,
socketaddrs: BTreeSet<SocketAddr>,
rt: &tokio::runtime::Handle,
) -> Result<DropGuard, IntoDiscoveryError> {
let spawn_rt = rt.clone();
let callback = move |node_id: &str, peer: &Peer| {
trace!(
node_id,
?peer,
"Received peer information from MdnsDiscovery"
);
let sender = sender.clone();
let node_id = node_id.to_string();
let peer = peer.clone();
spawn_rt.spawn(async move {
sender.send(Message::Discovery(node_id, peer)).await.ok();
});
};
let node_id_str = data_encoding::BASE32_NOPAD
.encode(node_id.as_bytes())
.to_ascii_lowercase();
let mut discoverer = Discoverer::new_interactive(N0_LOCAL_SWARM.to_string(), node_id_str)
.with_callback(callback)
.with_ip_class(IpClass::Auto);
if advertise {
let addrs = MdnsDiscovery::socketaddrs_to_addrs(&socketaddrs);
for addr in addrs {
discoverer = discoverer.with_addrs(addr.0, addr.1);
}
}
discoverer
.spawn(rt)
.map_err(|e| IntoDiscoveryError::from_err("mdns", e))
}
fn socketaddrs_to_addrs(socketaddrs: &BTreeSet<SocketAddr>) -> HashMap<u16, Vec<IpAddr>> {
let mut addrs: HashMap<u16, Vec<IpAddr>> = HashMap::default();
for socketaddr in socketaddrs {
addrs
.entry(socketaddr.port())
.and_modify(|a| a.push(socketaddr.ip()))
.or_insert(vec![socketaddr.ip()]);
}
addrs
}
}
fn peer_to_discovery_item(peer: &Peer, node_id: &NodeId) -> DiscoveryItem {
let direct_addresses: BTreeSet<SocketAddr> = peer
.addrs()
.iter()
.map(|(ip, port)| SocketAddr::new(*ip, *port))
.collect();
let user_data = if let Some(Some(user_data)) = peer.txt_attribute(USER_DATA_ATTRIBUTE) {
match user_data.parse() {
Err(err) => {
debug!("failed to parse user data from TXT attribute: {err}");
None
}
Ok(data) => Some(data),
}
} else {
None
};
let node_info = NodeInfo::new(*node_id)
.with_direct_addresses(direct_addresses)
.with_user_data(user_data);
DiscoveryItem::new(node_info, NAME, None)
}
impl Discovery for MdnsDiscovery {
fn resolve(&self, node_id: NodeId) -> Option<BoxStream<Result<DiscoveryItem, DiscoveryError>>> {
use futures_util::FutureExt;
let (send, recv) = mpsc::channel(20);
let discovery_sender = self.sender.clone();
let stream = async move {
discovery_sender
.send(Message::Resolve(node_id, send))
.await
.ok();
tokio_stream::wrappers::ReceiverStream::new(recv)
};
Some(Box::pin(stream.flatten_stream()))
}
fn publish(&self, data: &NodeData) {
if self.advertise {
self.local_addrs.set(Some(data.clone())).ok();
}
}
fn subscribe(&self) -> Option<BoxStream<DiscoveryEvent>> {
use futures_util::FutureExt;
let (sender, recv) = mpsc::channel(20);
let discovery_sender = self.sender.clone();
let stream = async move {
discovery_sender.send(Message::Subscribe(sender)).await.ok();
tokio_stream::wrappers::ReceiverStream::new(recv)
};
Some(Box::pin(stream.flatten_stream()))
}
}
#[cfg(test)]
mod tests {
mod run_in_isolation {
use iroh_base::SecretKey;
use n0_future::StreamExt;
use n0_snafu::{Error, Result, ResultExt};
use snafu::whatever;
use tracing_test::traced_test;
use super::super::*;
use crate::discovery::UserData;
#[tokio::test]
#[traced_test]
async fn mdns_publish_resolve() -> Result {
let (_, discovery_a) = make_discoverer(false)?;
let (node_id_b, discovery_b) = make_discoverer(true)?;
let user_data: UserData = "foobar".parse()?;
let node_data = NodeData::new(None, BTreeSet::from(["0.0.0.0:11111".parse().unwrap()]))
.with_user_data(Some(user_data.clone()));
let mut s1 = discovery_a
.subscribe()
.unwrap()
.filter(|event| match event {
DiscoveryEvent::Discovered(event) => event.node_id() == node_id_b,
_ => false,
});
let mut s2 = discovery_a
.subscribe()
.unwrap()
.filter(|event| match event {
DiscoveryEvent::Discovered(event) => event.node_id() == node_id_b,
_ => false,
});
tracing::debug!(?node_id_b, "Discovering node id b");
discovery_b.publish(&node_data);
let DiscoveryEvent::Discovered(s1_res) =
tokio::time::timeout(Duration::from_secs(5), s1.next())
.await
.context("timeout")?
.unwrap()
else {
panic!("Received unexpected discovery event");
};
let DiscoveryEvent::Discovered(s2_res) =
tokio::time::timeout(Duration::from_secs(5), s2.next())
.await
.context("timeout")?
.unwrap()
else {
panic!("Received unexpected discovery event");
};
assert_eq!(s1_res.node_info().data, node_data);
assert_eq!(s2_res.node_info().data, node_data);
Ok(())
}
#[tokio::test]
#[traced_test]
async fn mdns_publish_expire() -> Result {
let (_, discovery_a) = make_discoverer(false)?;
let (node_id_b, discovery_b) = make_discoverer(true)?;
let node_data = NodeData::new(None, BTreeSet::from(["0.0.0.0:11111".parse().unwrap()]))
.with_user_data(Some("".parse()?));
discovery_b.publish(&node_data);
let mut s1 = discovery_a.subscribe().unwrap();
tracing::debug!(?node_id_b, "Discovering node id b");
loop {
let event = tokio::time::timeout(Duration::from_secs(5), s1.next())
.await
.context("timeout")?
.expect("Stream should not be closed");
match event {
DiscoveryEvent::Discovered(item) if item.node_info().node_id == node_id_b => {
break;
}
_ => continue, }
}
drop(discovery_b);
tokio::time::sleep(Duration::from_secs(5)).await;
loop {
let event = tokio::time::timeout(Duration::from_secs(10), s1.next())
.await
.context("timeout waiting for expiration event")?
.expect("Stream should not be closed");
match event {
DiscoveryEvent::Expired(expired_node_id) if expired_node_id == node_id_b => {
break;
}
_ => continue, }
}
Ok(())
}
#[tokio::test]
#[traced_test]
async fn mdns_subscribe() -> Result {
let num_nodes = 5;
let mut node_ids = BTreeSet::new();
let mut discoverers = vec![];
let (_, discovery) = make_discoverer(false)?;
let node_data = NodeData::new(None, BTreeSet::from(["0.0.0.0:11111".parse().unwrap()]));
for i in 0..num_nodes {
let (node_id, discovery) = make_discoverer(true)?;
let user_data: UserData = format!("node{i}").parse()?;
let node_data = node_data.clone().with_user_data(Some(user_data.clone()));
node_ids.insert((node_id, Some(user_data)));
discovery.publish(&node_data);
discoverers.push(discovery);
}
let mut events = discovery.subscribe().unwrap();
let test = async move {
let mut got_ids = BTreeSet::new();
while got_ids.len() != num_nodes {
if let Some(DiscoveryEvent::Discovered(item)) = events.next().await {
if node_ids.contains(&(item.node_id(), item.user_data())) {
got_ids.insert((item.node_id(), item.user_data()));
}
} else {
whatever!(
"no more events, only got {} ids, expected {num_nodes}\n",
got_ids.len()
);
}
}
assert_eq!(got_ids, node_ids);
Ok::<_, Error>(())
};
tokio::time::timeout(Duration::from_secs(5), test)
.await
.context("timeout")?
}
#[tokio::test]
#[traced_test]
async fn non_advertising_node_not_discovered() -> Result {
let (_, discovery_a) = make_discoverer(false)?;
let (node_id_b, discovery_b) = make_discoverer(false)?;
let (node_id_c, discovery_c) = make_discoverer(true)?;
let node_data_c =
NodeData::new(None, BTreeSet::from(["0.0.0.0:22222".parse().unwrap()]));
discovery_c.publish(&node_data_c);
let node_data_b =
NodeData::new(None, BTreeSet::from(["0.0.0.0:11111".parse().unwrap()]));
discovery_b.publish(&node_data_b);
let mut stream_c = discovery_a.resolve(node_id_c).unwrap();
let result_c = tokio::time::timeout(Duration::from_secs(2), stream_c.next()).await;
assert!(result_c.is_ok(), "Advertising node should be discoverable");
let mut stream_b = discovery_a.resolve(node_id_b).unwrap();
let result_b = tokio::time::timeout(Duration::from_secs(2), stream_b.next()).await;
assert!(
result_b.is_err(),
"Expected timeout since node b isn't advertising"
);
Ok(())
}
fn make_discoverer(advertise: bool) -> Result<(PublicKey, MdnsDiscovery)> {
let node_id = SecretKey::generate(rand::thread_rng()).public();
Ok((node_id, MdnsDiscovery::new(node_id, advertise)?))
}
}
}