use std::future::Future;
use agnostic_net::Net;
use async_channel::Sender;
use hick_udp::{
MulticastOptionsV4, MulticastOptionsV6, try_bind_v4, try_bind_v6, try_join_v4, try_join_v6,
};
use mdns_proto::{QuerySpec, ServiceSpec};
use hick_trace::*;
use crate::{
command::{Command, QueryStarted, ServiceRegistered},
driver::{self, BoundSockets},
error::{RegisterError, ServerError, StartQueryError},
options::ServerOptions,
query::Query,
service::Service,
};
#[derive(Clone)]
pub struct Endpoint {
cmd: Sender<Command>,
#[cfg(feature = "stats")]
stats: std::sync::Arc<stats::Stats>,
}
impl Endpoint {
pub async fn server<N: Net>(opts: ServerOptions) -> Result<Self, ServerError> {
if !opts.ipv4() && !opts.ipv6() {
return Err(ServerError::NoFamilyEnabled);
}
let interface_index = match opts.interface_index() {
Some(i) => i,
None => pick_default_interface_index(opts.ipv4(), opts.ipv6()).ok_or_else(|| {
ServerError::Io(std::io::Error::new(
std::io::ErrorKind::NotFound,
"no multicast-capable interface found",
))
})?,
};
let iface_has_v4 = match getifs::interface_by_index(interface_index) {
Ok(Some(i)) => matches!(i.ipv4_addrs(), Ok(ref a) if !a.is_empty()),
_ => false,
};
let iface_has_v6 = match getifs::interface_by_index(interface_index) {
Ok(Some(i)) => matches!(i.ipv6_addrs(), Ok(ref a) if !a.is_empty()),
_ => false,
};
let bind_v4 = opts.ipv4() && iface_has_v4;
let bind_v6 = opts.ipv6() && iface_has_v6;
if !bind_v4 && !bind_v6 {
return Err(ServerError::Io(std::io::Error::new(
std::io::ErrorKind::AddrNotAvailable,
"interface has no address in any requested family",
)));
}
let v4 = if bind_v4 {
match try_bind_v4(MulticastOptionsV4::new(interface_index)) {
Ok(std_sock) => {
debug!(interface_index, "bound v4 mDNS socket");
match try_join_v4(&std_sock, interface_index) {
Ok(()) => {
debug!(interface_index, "joined v4 mDNS multicast group");
}
Err(e) => {
warn!(error = %e, interface_index, "failed to join v4 mDNS multicast group");
return Err(map_join_to_bind_v4(e));
}
}
std_sock.set_nonblocking(true)?;
let async_sock = N::UdpSocket::try_from(std_sock).map_err(ServerError::WrapSocket)?;
Some(async_sock)
}
Err(e) => {
warn!(error = %e, interface_index, "failed to bind v4 mDNS socket");
return Err(ServerError::BindV4(e));
}
}
} else {
None
};
let v6 = if bind_v6 {
match try_bind_v6(MulticastOptionsV6::new(interface_index)) {
Ok(std_sock) => {
debug!(interface_index, "bound v6 mDNS socket");
match try_join_v6(&std_sock, interface_index) {
Ok(()) => {
debug!(interface_index, "joined v6 mDNS multicast group");
}
Err(e) => {
warn!(error = %e, interface_index, "failed to join v6 mDNS multicast group");
return Err(map_join_to_bind_v6(e));
}
}
std_sock.set_nonblocking(true)?;
let async_sock = N::UdpSocket::try_from(std_sock).map_err(ServerError::WrapSocket)?;
Some(async_sock)
}
Err(e) => {
warn!(error = %e, interface_index, "failed to bind v6 mDNS socket");
return Err(ServerError::BindV6(e));
}
}
} else {
None
};
let (cmd_tx, cmd_rx) = async_channel::unbounded::<Command>();
let sockets = BoundSockets {
v4,
v6,
interface_index,
};
#[cfg(feature = "stats")]
let mut stats_slot: Option<std::sync::Arc<stats::Stats>> = None;
driver::spawn::<N>(
opts,
sockets,
cmd_rx,
#[cfg(feature = "stats")]
&mut stats_slot,
);
Ok(Self {
cmd: cmd_tx,
#[cfg(feature = "stats")]
stats: stats_slot.expect("spawn always populates stats_slot when stats feature is enabled"),
})
}
#[cfg(feature = "stats")]
#[cfg_attr(docsrs, doc(cfg(feature = "stats")))]
pub fn stats(&self) -> stats::StatsSnapshot {
self.stats.snapshot()
}
pub(crate) fn spawn_lookup<F>(&self, fut: F) -> Result<(), StartQueryError>
where
F: Future<Output = ()> + Send + 'static,
{
self
.cmd
.try_send(Command::SpawnLookup {
task: Box::pin(fut),
})
.map_err(|_| StartQueryError::DriverGone)
}
pub async fn register_service(&self, spec: ServiceSpec) -> Result<Service, RegisterError> {
let (reply_tx, reply_rx) = futures::channel::oneshot::channel();
self
.cmd
.send(Command::RegisterService {
spec,
reply: reply_tx,
})
.await
.map_err(|_| RegisterError::DriverGone)?;
let ServiceRegistered {
handle,
mailbox,
doorbell,
} = reply_rx.await.map_err(|_| RegisterError::DriverGone)??;
Ok(Service::new(handle, mailbox, doorbell, self.cmd.clone()))
}
pub async fn start_query(&self, spec: QuerySpec) -> Result<Query, StartQueryError> {
let (reply_tx, reply_rx) = futures::channel::oneshot::channel();
self
.cmd
.send(Command::StartQuery {
spec,
reply: reply_tx,
})
.await
.map_err(|_| StartQueryError::DriverGone)?;
let QueryStarted {
handle,
mailbox,
doorbell,
} = reply_rx.await.map_err(|_| StartQueryError::DriverGone)??;
Ok(Query::new(handle, mailbox, doorbell, self.cmd.clone()))
}
}
fn map_join_to_bind_v4(e: hick_udp::JoinError) -> ServerError {
match e {
hick_udp::JoinError::Io(io) => ServerError::BindV4(hick_udp::BindError::Io(io)),
hick_udp::JoinError::InterfaceNotFound(d) => {
ServerError::BindV4(hick_udp::BindError::InterfaceNotFound(d))
}
_ => ServerError::Io(std::io::Error::other("unknown JoinError variant")),
}
}
fn map_join_to_bind_v6(e: hick_udp::JoinError) -> ServerError {
match e {
hick_udp::JoinError::Io(io) => ServerError::BindV6(hick_udp::BindError::Io(io)),
hick_udp::JoinError::InterfaceNotFound(d) => {
ServerError::BindV6(hick_udp::BindError::InterfaceNotFound(d))
}
_ => ServerError::Io(std::io::Error::other("unknown JoinError variant")),
}
}
fn pick_default_interface_index(want_v4: bool, want_v6: bool) -> Option<u32> {
let ifs = getifs::interfaces().ok()?;
let has_v4 = |i: &getifs::Interface| matches!(i.ipv4_addrs(), Ok(ref v) if !v.is_empty());
let has_v6 = |i: &getifs::Interface| matches!(i.ipv6_addrs(), Ok(ref v) if !v.is_empty());
let multicast_up_non_loopback = |i: &getifs::Interface| -> bool {
let f = i.flags();
f.contains(getifs::Flags::UP)
&& f.contains(getifs::Flags::MULTICAST)
&& !f.contains(getifs::Flags::LOOPBACK)
&& i.index() != 0
};
let loopback_up = |i: &getifs::Interface| -> bool {
i.flags().contains(getifs::Flags::LOOPBACK) && i.flags().contains(getifs::Flags::UP)
};
let strict =
|i: &&getifs::Interface| -> bool { (!want_v4 || has_v4(i)) && (!want_v6 || has_v6(i)) };
let loose = |i: &&getifs::Interface| -> bool { (want_v4 && has_v4(i)) || (want_v6 && has_v6(i)) };
let strict_non_loopback = ifs
.iter()
.find(|i| multicast_up_non_loopback(i) && strict(i));
let loose_non_loopback = ifs
.iter()
.find(|i| multicast_up_non_loopback(i) && loose(i));
let strict_loopback = ifs.iter().find(|i| loopback_up(i) && strict(i));
let loose_loopback = ifs.iter().find(|i| loopback_up(i) && loose(i));
strict_non_loopback
.or(loose_non_loopback)
.or(strict_loopback)
.or(loose_loopback)
.map(|i| i.index())
}
#[cfg(test)]
mod tests;