use std::{
collections::{HashMap, HashSet, VecDeque},
net::{IpAddr, Ipv4Addr, Ipv6Addr},
sync::{Arc, Mutex, MutexGuard},
time::{Duration, Instant},
};
use async_channel::{Receiver, Sender};
use bytes::Bytes;
use futures::{
FutureExt, StreamExt, pin_mut, select_biased,
stream::{BoxStream, SelectAll},
};
use mdns_proto::{
Name, QuerySpec,
wire::{A, AAAA, NameRef, Ptr, ResourceType, Srv, Txt},
};
use smol_str::SmolStr;
use crate::{
Endpoint, QueryEvent,
error::StartQueryError,
query::{DroppedHandle, Query},
};
pub const DEFAULT_MAX_ENTRIES: usize = 64;
const MAX_ADDRS_PER_HOST: usize = 16;
#[derive(Debug, Clone)]
pub struct ServiceEntry {
instance: Name,
host: Name,
port: u16,
ipv4: Arc<[Ipv4Addr]>,
ipv6: Arc<[Ipv6Addr]>,
txt: Arc<[Bytes]>,
}
impl ServiceEntry {
#[inline]
pub fn instance_name(&self) -> &Name {
&self.instance
}
#[inline]
pub fn host(&self) -> &Name {
&self.host
}
#[inline]
pub const fn port(&self) -> u16 {
self.port
}
#[inline]
pub fn ipv4_addresses(&self) -> &[Ipv4Addr] {
&self.ipv4
}
#[inline]
pub fn ipv6_addresses(&self) -> &[Ipv6Addr] {
&self.ipv6
}
pub fn addresses(&self) -> impl Iterator<Item = IpAddr> + '_ {
self
.ipv4
.iter()
.copied()
.map(IpAddr::V4)
.chain(self.ipv6.iter().copied().map(IpAddr::V6))
}
#[inline]
pub fn txt(&self) -> &[Bytes] {
&self.txt
}
}
#[derive(Debug, Clone)]
pub struct QueryParam {
service: Name,
timeout: Duration,
resolve_timeout: Option<Duration>,
unicast_response: bool,
max_entries: usize,
}
impl QueryParam {
pub fn new(service: Name) -> Self {
Self {
service,
timeout: Duration::from_secs(1),
resolve_timeout: None,
unicast_response: false,
max_entries: DEFAULT_MAX_ENTRIES,
}
}
#[must_use]
pub const fn with_timeout(mut self, timeout: Duration) -> Self {
self.timeout = timeout;
self
}
#[must_use]
pub const fn with_resolve_timeout(mut self, timeout: Duration) -> Self {
self.resolve_timeout = Some(timeout);
self
}
#[must_use]
pub const fn with_unicast_response(mut self, unicast: bool) -> Self {
self.unicast_response = unicast;
self
}
#[must_use]
pub const fn with_max_entries(mut self, max: usize) -> Self {
self.max_entries = if max == 0 { 1 } else { max };
self
}
}
#[derive(Clone)]
#[allow(clippy::upper_case_acronyms)] enum Step {
Ptr,
Srv(SmolStr),
Txt(SmolStr),
A(SmolStr),
AAAA(SmolStr),
}
struct Tagged {
step: Step,
event: QueryEvent,
}
#[derive(Clone)]
struct Start {
name: Name,
step: Step,
}
impl Start {
const fn qtype(&self) -> ResourceType {
match self.step {
Step::Srv(_) => ResourceType::Srv,
Step::Txt(_) => ResourceType::Txt,
Step::A(_) => ResourceType::A,
Step::AAAA(_) => ResourceType::AAAA,
Step::Ptr => ResourceType::Ptr,
}
}
}
struct Builder {
instance: Name,
host: Option<Name>,
host_key: Option<SmolStr>,
has_srv: bool,
port: u16,
ipv4: Vec<Ipv4Addr>,
ipv6: Vec<Ipv6Addr>,
txt: Option<Vec<Bytes>>,
emitted: bool,
}
impl Builder {
fn new(instance: Name) -> Self {
Self {
instance,
host: None,
host_key: None,
has_srv: false,
port: 0,
ipv4: Vec::new(),
ipv6: Vec::new(),
txt: None,
emitted: false,
}
}
fn complete(&self) -> bool {
self.has_srv && self.txt.is_some() && !(self.ipv4.is_empty() && self.ipv6.is_empty())
}
fn finalize(&self) -> Option<ServiceEntry> {
Some(ServiceEntry {
instance: self.instance.clone(),
host: self.host.clone()?,
port: self.port,
ipv4: self.ipv4.as_slice().into(),
ipv6: self.ipv6.as_slice().into(),
txt: self.txt.as_deref()?.into(),
})
}
}
#[derive(Default)]
struct HostAddrs {
ipv4: Vec<Ipv4Addr>,
ipv6: Vec<Ipv6Addr>,
}
struct Resolver {
builders: HashMap<SmolStr, Builder>,
host_addrs: HashMap<SmolStr, HostAddrs>,
hosts_queried: HashSet<SmolStr>,
ready: VecDeque<ServiceEntry>,
max_entries: usize,
max_hosts: usize,
dropped: u64,
}
impl Resolver {
fn new(max_entries: usize) -> Self {
Self {
builders: HashMap::new(),
host_addrs: HashMap::new(),
hosts_queried: HashSet::new(),
ready: VecDeque::new(),
max_entries,
max_hosts: max_entries,
dropped: 0,
}
}
fn on_ptr(&mut self, instance: Name) -> Vec<Start> {
let key = fold(&instance);
if self.builders.contains_key(&key) {
return Vec::new(); }
if self.builders.len() >= self.max_entries {
self.dropped = self.dropped.saturating_add(1);
return Vec::new();
}
self
.builders
.insert(key.clone(), Builder::new(instance.clone()));
vec![
Start {
name: instance.clone(),
step: Step::Srv(key.clone()),
},
Start {
name: instance,
step: Step::Txt(key),
},
]
}
fn on_srv(&mut self, inst_key: &str, host: Name, port: u16) -> Vec<Start> {
let host_key = fold(&host);
let cached = self.host_addrs.get(&host_key);
let mut changed = false;
if let Some(b) = self.builders.get_mut(inst_key) {
let host_changed = b.host_key.as_deref() != Some(host_key.as_str());
if host_changed {
b.ipv4.clear();
b.ipv6.clear();
}
changed = host_changed || b.port != port;
b.has_srv = true;
b.host = Some(host.clone());
b.host_key = Some(host_key.clone());
b.port = port;
if let Some(addrs) = cached {
for &a in &addrs.ipv4 {
push_capped(&mut b.ipv4, a);
}
for &a in &addrs.ipv6 {
push_capped(&mut b.ipv6, a);
}
}
}
self.try_emit(inst_key, changed);
if self.hosts_queried.contains(&host_key) {
return Vec::new(); }
if self.hosts_queried.len() >= self.max_hosts {
self.dropped = self.dropped.saturating_add(1);
return Vec::new();
}
self.hosts_queried.insert(host_key.clone());
vec![
Start {
name: host.clone(),
step: Step::A(host_key.clone()),
},
Start {
name: host,
step: Step::AAAA(host_key),
},
]
}
fn on_txt(&mut self, inst_key: &str, segs: Vec<Bytes>) {
let mut changed = false;
if let Some(b) = self.builders.get_mut(inst_key) {
changed = b.txt.as_deref() != Some(segs.as_slice());
b.txt = Some(segs);
}
self.try_emit(inst_key, changed);
}
fn on_addr(&mut self, host_key: &str, addr: IpAddr) {
let cache = self.host_addrs.entry(SmolStr::from(host_key)).or_default();
match addr {
IpAddr::V4(a) => push_capped(&mut cache.ipv4, a),
IpAddr::V6(a) => push_capped(&mut cache.ipv6, a),
};
let keys: Vec<SmolStr> = self
.builders
.iter()
.filter(|(_, b)| b.host_key.as_deref() == Some(host_key))
.map(|(k, _)| k.clone())
.collect();
for k in keys {
let added = match self.builders.get_mut(&k) {
Some(b) => match addr {
IpAddr::V4(a) => push_capped(&mut b.ipv4, a),
IpAddr::V6(a) => push_capped(&mut b.ipv6, a),
},
None => false,
};
if added {
self.try_emit(&k, true);
}
}
}
fn try_emit(&mut self, inst_key: &str, allow_reemit: bool) {
if let Some(b) = self.builders.get_mut(inst_key) {
if !b.complete() {
return;
}
if !b.emitted {
if let Some(entry) = b.finalize() {
b.emitted = true;
self.ready.push_back(entry);
}
} else if allow_reemit && let Some(entry) = b.finalize() {
self.ready.push_back(entry);
}
}
}
fn take_ready(&mut self) -> Option<ServiceEntry> {
self.ready.pop_front()
}
}
struct LookupQueue {
ready: VecDeque<ServiceEntry>,
done: bool,
dropped: u64,
}
impl LookupQueue {
#[inline(always)]
const fn new() -> Self {
Self {
ready: VecDeque::new(),
done: false,
dropped: 0,
}
}
fn enqueue(&mut self, entry: ServiceEntry) {
if let Some(slot) = self
.ready
.iter_mut()
.find(|e| e.instance.as_str() == entry.instance.as_str())
{
*slot = entry;
} else {
self.ready.push_back(entry);
}
}
}
fn lock(q: &Mutex<LookupQueue>) -> MutexGuard<'_, LookupQueue> {
q.lock().unwrap_or_else(|poisoned| poisoned.into_inner())
}
struct LookupDriver {
endpoint: Endpoint,
streams: SelectAll<BoxStream<'static, Tagged>>,
resolver: Resolver,
ptr_drops: DroppedHandle,
queue: Arc<Mutex<LookupQueue>>,
doorbell: Sender<()>,
cancel: Receiver<()>,
resolve_timeout: Duration,
unicast: bool,
}
impl LookupDriver {
async fn run(mut self) {
loop {
let tagged = {
let next = self.streams.next().fuse();
let stop = self.cancel.recv().fuse();
pin_mut!(next, stop);
select_biased! {
_ = stop => None, t = next => t, }
};
match tagged {
Some(t) => self.process(t).await,
None => break,
}
}
{
let mut q = lock(&self.queue);
q.done = true;
q.dropped = self.dropped_total();
}
let _ = self.doorbell.try_send(());
}
async fn process(&mut self, tagged: Tagged) {
for start in feed(&mut self.resolver, tagged) {
self.launch(start).await;
}
self.flush();
}
fn dropped_total(&self) -> u64 {
self.resolver.dropped.saturating_add(self.ptr_drops.get())
}
fn flush(&mut self) {
let mut woke = false;
let mut q = lock(&self.queue);
while let Some(entry) = self.resolver.take_ready() {
q.enqueue(entry);
woke = true;
}
q.dropped = self.dropped_total();
drop(q);
if woke {
let _ = self.doorbell.try_send(());
}
}
async fn launch(&mut self, start: Start) {
let qtype = start.qtype();
let spec = QuerySpec::new(start.name, qtype)
.with_timeout(self.resolve_timeout)
.with_unicast_response(self.unicast);
if let Ok(query) = self.endpoint.start_query(spec).await {
self.streams.push(tagged_stream(query, start.step));
}
}
}
pub struct Lookup {
queue: Arc<Mutex<LookupQueue>>,
doorbell: async_channel::Receiver<()>,
_cancel: async_channel::Sender<()>,
}
impl Lookup {
pub async fn next(&mut self) -> Option<ServiceEntry> {
loop {
{
let mut q = lock(&self.queue);
if let Some(entry) = q.ready.pop_front() {
return Some(entry);
}
if q.done {
return None;
}
}
if self.doorbell.recv().await.is_err() {
return lock(&self.queue).ready.pop_front();
}
}
}
pub fn dropped(&self) -> u64 {
lock(&self.queue).dropped
}
}
impl Endpoint {
pub async fn browse(&self, param: QueryParam) -> Result<Lookup, StartQueryError> {
let resolve_timeout = param.resolve_timeout.unwrap_or(param.timeout);
let ptr_spec = QuerySpec::new(param.service, ResourceType::Ptr)
.with_timeout(param.timeout)
.with_unicast_response(param.unicast_response)
.with_max_answers(param.max_entries);
let ptr_query = self.start_query(ptr_spec).await?;
let ptr_drops = ptr_query.dropped_handle();
let mut streams = SelectAll::new();
streams.push(tagged_stream(ptr_query, Step::Ptr));
let queue = Arc::new(Mutex::new(LookupQueue::new()));
let (doorbell_tx, doorbell_rx) = async_channel::bounded(1);
let (cancel_tx, cancel_rx) = async_channel::bounded(1);
let driver = LookupDriver {
endpoint: self.clone(),
streams,
resolver: Resolver::new(param.max_entries),
ptr_drops,
queue: Arc::clone(&queue),
doorbell: doorbell_tx,
cancel: cancel_rx,
resolve_timeout,
unicast: param.unicast_response,
};
self.spawn_lookup(driver.run())?;
Ok(Lookup {
queue,
doorbell: doorbell_rx,
_cancel: cancel_tx,
})
}
pub async fn lookup(&self, service: Name, timeout: Duration) -> Result<Lookup, StartQueryError> {
self
.browse(QueryParam::new(service).with_timeout(timeout))
.await
}
pub async fn resolve_host(
&self,
host: Name,
timeout: Duration,
) -> Result<Vec<IpAddr>, StartQueryError> {
let host_key = fold(&host);
let a = self
.start_query(QuerySpec::new(host.clone(), ResourceType::A).with_timeout(timeout))
.await?;
let aaaa = self
.start_query(QuerySpec::new(host, ResourceType::AAAA).with_timeout(timeout))
.await?;
let mut streams = SelectAll::new();
streams.push(tagged_stream(a, Step::A(host_key.clone())));
streams.push(tagged_stream(aaaa, Step::AAAA(host_key)));
let mut ipv4: Vec<Ipv4Addr> = Vec::new();
let mut ipv6: Vec<Ipv6Addr> = Vec::new();
while let Some(tagged) = streams.next().await {
let ans = match tagged.event {
QueryEvent::Answer(a) => a,
QueryEvent::Terminal(_) => continue,
};
match ans.rtype() {
ResourceType::A => {
if let Ok(r) = A::try_from_rdata(ans.rdata_slice()) {
push_capped(&mut ipv4, r.addr());
}
}
ResourceType::AAAA => {
if let Ok(r) = AAAA::try_from_rdata(ans.rdata_slice()) {
push_capped(&mut ipv6, r.addr());
}
}
_ => {}
}
}
Ok(
ipv4
.into_iter()
.map(IpAddr::V4)
.chain(ipv6.into_iter().map(IpAddr::V6))
.collect(),
)
}
pub async fn resolve_instance(
&self,
instance: Name,
timeout: Duration,
) -> Result<Option<ServiceEntry>, StartQueryError> {
let deadline = Instant::now().checked_add(timeout);
let remaining = || deadline.map_or(timeout, |d| d.saturating_duration_since(Instant::now()));
let mut resolver = Resolver::new(1);
let mut streams = SelectAll::new();
for start in resolver.on_ptr(instance) {
streams.push(self.launch_resolve(start, remaining()).await?);
}
while let Some(tagged) = streams.next().await {
for start in feed(&mut resolver, tagged) {
streams.push(self.launch_resolve(start, remaining()).await?);
}
if let Some(entry) = resolver.take_ready() {
return Ok(Some(entry));
}
}
Ok(resolver.take_ready())
}
async fn launch_resolve(
&self,
start: Start,
timeout: Duration,
) -> Result<futures::stream::BoxStream<'static, Tagged>, StartQueryError> {
let qtype = start.qtype();
let query = self
.start_query(QuerySpec::new(start.name, qtype).with_timeout(timeout))
.await?;
Ok(tagged_stream(query, start.step))
}
}
fn feed(resolver: &mut Resolver, tagged: Tagged) -> Vec<Start> {
let answer = match tagged.event {
QueryEvent::Answer(a) => a,
QueryEvent::Terminal(_) => return Vec::new(),
};
match tagged.step {
Step::Ptr => {
if answer.rtype() != ResourceType::Ptr {
return Vec::new();
}
match parse_name(answer.rdata_slice()) {
Some(instance) => resolver.on_ptr(instance),
None => Vec::new(),
}
}
Step::Srv(inst_key) => {
if answer.rtype() != ResourceType::Srv {
return Vec::new();
}
match parse_srv(answer.rdata_slice()) {
Some((host, port)) => resolver.on_srv(&inst_key, host, port),
None => Vec::new(),
}
}
Step::Txt(inst_key) => {
if answer.rtype() != ResourceType::Txt {
return Vec::new();
}
resolver.on_txt(&inst_key, parse_txt(answer.rdata_slice()));
Vec::new()
}
Step::A(host_key) => {
if let Ok(r) = A::try_from_rdata(answer.rdata_slice())
&& answer.rtype() == ResourceType::A
{
resolver.on_addr(&host_key, IpAddr::V4(r.addr()));
}
Vec::new()
}
Step::AAAA(host_key) => {
if let Ok(r) = AAAA::try_from_rdata(answer.rdata_slice())
&& answer.rtype() == ResourceType::AAAA
{
resolver.on_addr(&host_key, IpAddr::V6(r.addr()));
}
Vec::new()
}
}
}
fn tagged_stream(query: Query, step: Step) -> futures::stream::BoxStream<'static, Tagged> {
futures::stream::unfold((query, step), |(mut query, step)| async move {
let event = query.next().await?;
let tagged = Tagged {
step: step.clone(),
event,
};
Some((tagged, (query, step)))
})
.boxed()
}
fn push_capped<T: PartialEq>(v: &mut Vec<T>, item: T) -> bool {
if v.len() >= MAX_ADDRS_PER_HOST || v.contains(&item) {
return false;
}
v.push(item);
true
}
fn fold(name: &Name) -> SmolStr {
SmolStr::new(name.as_str())
}
fn name_from_ref(nr: &NameRef<'_>) -> Option<Name> {
let mut s = String::new();
for label in nr.labels() {
let label = label.ok()?;
if label.is_empty() {
break; }
if label.iter().any(|&b| b >= 0x80 || b == b'.') {
return None; }
s.push_str(core::str::from_utf8(label).ok()?);
s.push('.');
}
if s.is_empty() {
return None;
}
Name::try_from_str(&s).ok()
}
fn parse_name(rdata: &[u8]) -> Option<Name> {
let ptr = Ptr::try_from_message(rdata, 0, rdata.len()).ok()?;
name_from_ref(ptr.target())
}
fn parse_srv(rdata: &[u8]) -> Option<(Name, u16)> {
let srv = Srv::try_from_message(rdata, 0, rdata.len()).ok()?;
let host = name_from_ref(srv.target())?;
Some((host, srv.port()))
}
fn parse_txt(rdata: &[u8]) -> Vec<Bytes> {
Txt::from_rdata(rdata)
.segments()
.map_while(Result::ok)
.map(Bytes::copy_from_slice)
.collect()
}
#[cfg(test)]
#[allow(clippy::unwrap_used)]
mod tests;