use std::collections::HashSet;
use std::ffi::c_void;
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr};
use std::num::NonZeroU32;
use std::sync::{Arc, Mutex, Weak};
use log::trace;
use tokio::runtime::Handle;
use tokio::sync::mpsc::unbounded_channel;
use widestring::U16CString;
use windows::Win32::NetworkManagement::Dns;
use windows::core::{PCWSTR, PWSTR};
use crate::browse::{
BrowseEvent, BrowseEventReceiver, BrowseEventSender, DiscoveredService, ServiceBrowseError,
TxtRecord, trim_dot,
};
const META_QUERY_TYPE: &str = "_services._dns-sd._udp.local";
const DEFAULT_DOMAIN: &str = "local";
const DNS_REQUEST_PENDING: i32 = 9506;
type Registry = Arc<Mutex<Vec<BrowseEntry>>>;
pub(crate) struct BrowseGuard {
_registry: Registry,
}
struct BrowseEntry {
_cancel: CancelHandle,
_context: Box<BrowseContext>,
_query_name: U16CString,
}
struct CancelHandle(Dns::DNS_SERVICE_CANCEL);
unsafe impl Send for CancelHandle {}
impl Drop for CancelHandle {
fn drop(&mut self) {
unsafe {
let _ = Dns::DnsServiceBrowseCancel(&self.0);
}
}
}
struct BrowseContext {
tx: BrowseEventSender,
handle: Handle,
registry: Weak<Mutex<Vec<BrowseEntry>>>,
is_meta: bool,
interface: u32,
service_type: String,
domain: String,
seen: Mutex<HashSet<String>>,
}
unsafe impl Send for BrowseContext {}
unsafe impl Sync for BrowseContext {}
pub(crate) async fn browse_start(
service_type: &Option<String>,
domain: &Option<String>,
interface_index: Option<NonZeroU32>,
) -> Result<(BrowseEventReceiver, BrowseGuard), ServiceBrowseError> {
let (tx, rx) = unbounded_channel();
let handle = Handle::current();
let interface = interface_index.map(|i| i.get()).unwrap_or(0); let domain = domain.clone().unwrap_or_else(|| DEFAULT_DOMAIN.to_string());
let registry: Registry = Arc::new(Mutex::new(Vec::new()));
let (query_name, service_type, is_meta) = match service_type {
Some(service_type) => (
format!("{service_type}.{domain}"),
service_type.clone(),
false,
),
None => (
META_QUERY_TYPE.to_string(),
META_QUERY_TYPE.to_string(),
true,
),
};
start_browse(
query_name,
service_type,
domain,
interface,
is_meta,
tx,
handle,
®istry,
)
.map_err(ServiceBrowseError::BrowseFailed)?;
Ok((
rx,
BrowseGuard {
_registry: registry,
},
))
}
#[allow(clippy::too_many_arguments)]
fn start_browse(
query_name: String,
service_type: String,
domain: String,
interface: u32,
is_meta: bool,
tx: BrowseEventSender,
handle: Handle,
registry: &Registry,
) -> Result<(), String> {
let query_name_w = U16CString::from_str(&query_name).map_err(|e| e.to_string())?;
let context = Box::new(BrowseContext {
tx,
handle,
registry: Arc::downgrade(registry),
is_meta,
interface,
service_type,
domain,
seen: Mutex::new(HashSet::new()),
});
let context_ptr = &*context as *const BrowseContext as *mut c_void;
let request = Dns::DNS_SERVICE_BROWSE_REQUEST {
Version: Dns::DNS_QUERY_REQUEST_VERSION1.0,
InterfaceIndex: interface,
QueryName: PCWSTR(query_name_w.as_ptr()),
Anonymous: Dns::DNS_SERVICE_BROWSE_REQUEST_0 {
pBrowseCallback: Some(browse_callback),
},
pQueryContext: context_ptr,
};
let mut cancel = Dns::DNS_SERVICE_CANCEL::default();
let result = unsafe { Dns::DnsServiceBrowse(&request, &mut cancel) };
if result != DNS_REQUEST_PENDING {
return Err(format!("DnsServiceBrowse failed with status {result}"));
}
let entry = BrowseEntry {
_cancel: CancelHandle(cancel),
_context: context,
_query_name: query_name_w,
};
registry
.lock()
.map_err(|_| "browse registry poisoned".to_string())?
.push(entry);
Ok(())
}
unsafe extern "system" fn browse_callback(
status: u32,
context: *const c_void,
records: *const Dns::DNS_RECORDW,
) {
let ctx = unsafe { &*(context as *const BrowseContext) };
if status != 0 {
let _ = ctx.tx.send(Err(ServiceBrowseError::BrowseFailed(format!(
"browse callback status {status}"
))));
}
let mut record = records;
while !record.is_null() {
let r = unsafe { &*record };
if r.wType == Dns::DNS_TYPE_PTR.0 {
let target = trim_dot(&unsafe { pwstr_to_string(r.Data.Ptr.pNameHost) });
handle_ptr(ctx, target);
}
record = r.pNext;
}
if !records.is_null() {
unsafe {
Dns::DnsFree(Some(records as *const c_void), Dns::DnsFreeRecordList);
}
}
}
fn handle_ptr(ctx: &BrowseContext, target: String) {
{
let mut seen = match ctx.seen.lock() {
Ok(seen) => seen,
Err(_) => return,
};
if !seen.insert(target.clone()) {
return; }
}
if ctx.is_meta {
let (service_type, domain) = match target.rsplit_once('.') {
Some((service_type, domain)) => (service_type.to_string(), domain.to_string()),
None => (target.clone(), ctx.domain.clone()),
};
trace!("discovered service type {service_type:?} in domain {domain:?}");
if let Some(registry) = ctx.registry.upgrade() {
let _ = start_browse(
target,
service_type,
domain,
ctx.interface,
false,
ctx.tx.clone(),
ctx.handle.clone(),
®istry,
);
}
} else {
let label = instance_label(&target, &ctx.service_type, &ctx.domain);
start_resolve(ctx, target, label);
}
}
struct ResolveContext {
tx: BrowseEventSender,
name: String,
service_type: String,
domain: String,
_query_name: U16CString,
}
fn start_resolve(ctx: &BrowseContext, full_name: String, label: String) {
let query_name_w = match U16CString::from_str(&full_name) {
Ok(query_name_w) => query_name_w,
Err(_) => return,
};
let resolve_ctx = Box::new(ResolveContext {
tx: ctx.tx.clone(),
name: label,
service_type: ctx.service_type.clone(),
domain: ctx.domain.clone(),
_query_name: query_name_w,
});
let query_ptr = resolve_ctx._query_name.as_ptr();
let resolve_ctx_ptr = Box::into_raw(resolve_ctx);
let request = Dns::DNS_SERVICE_RESOLVE_REQUEST {
Version: Dns::DNS_QUERY_REQUEST_VERSION1.0,
InterfaceIndex: ctx.interface,
QueryName: PWSTR(query_ptr as *mut u16),
pResolveCompletionCallback: Some(resolve_callback),
pQueryContext: resolve_ctx_ptr as *mut c_void,
};
let mut cancel = Dns::DNS_SERVICE_CANCEL::default();
let result = unsafe { Dns::DnsServiceResolve(&request, &mut cancel) };
if result != DNS_REQUEST_PENDING {
let resolve_ctx = unsafe { Box::from_raw(resolve_ctx_ptr) };
let _ = resolve_ctx.tx.send(Err(ServiceBrowseError::ResolveFailed(
resolve_ctx.name.clone(),
format!("DnsServiceResolve failed with status {result}"),
)));
}
}
unsafe extern "system" fn resolve_callback(
status: u32,
context: *const c_void,
instance: *const Dns::DNS_SERVICE_INSTANCE,
) {
let resolve_ctx = unsafe { Box::from_raw(context as *mut ResolveContext) };
if status != 0 || instance.is_null() {
let _ = resolve_ctx.tx.send(Err(ServiceBrowseError::ResolveFailed(
resolve_ctx.name.clone(),
format!("resolve callback status {status}"),
)));
if !instance.is_null() {
unsafe { Dns::DnsServiceFreeInstance(instance) };
}
return;
}
let inst = unsafe { &*instance };
let mut addresses = Vec::new();
if !inst.ip4Address.is_null() {
let raw = unsafe { *inst.ip4Address };
addresses.push(IpAddr::V4(Ipv4Addr::from(u32::from_be(raw))));
}
if !inst.ip6Address.is_null() {
let bytes = unsafe { (*inst.ip6Address).IP6Byte };
addresses.push(IpAddr::V6(Ipv6Addr::from(bytes)));
}
let txt_records = unsafe { read_txt(inst) };
let service = DiscoveredService {
name: resolve_ctx.name.clone(),
service_type: resolve_ctx.service_type.clone(),
domain: resolve_ctx.domain.clone(),
host_name: trim_dot(&unsafe { pwstr_to_string(inst.pszHostName) }),
port: inst.wPort,
addresses,
txt_records,
interface_index: NonZeroU32::new(inst.dwInterfaceIndex),
};
let _ = resolve_ctx.tx.send(Ok(BrowseEvent::Found(service)));
unsafe { Dns::DnsServiceFreeInstance(instance) };
}
unsafe fn read_txt(inst: &Dns::DNS_SERVICE_INSTANCE) -> Vec<TxtRecord> {
let count = inst.dwPropertyCount as usize;
if count == 0 || inst.keys.is_null() {
return Vec::new();
}
let keys = unsafe { std::slice::from_raw_parts(inst.keys, count) };
let values: &[PWSTR] = if inst.values.is_null() {
&[]
} else {
unsafe { std::slice::from_raw_parts(inst.values, count) }
};
let mut records = Vec::with_capacity(count);
for (i, key_ptr) in keys.iter().enumerate() {
let key = unsafe { pwstr_to_string(*key_ptr) };
if key.is_empty() {
continue;
}
let value = match values.get(i) {
Some(v) if !v.0.is_null() => Some(unsafe { pwstr_to_string(*v) }.into_bytes()),
_ => None,
};
records.push(TxtRecord { key, value });
}
records
}
unsafe fn pwstr_to_string(p: PWSTR) -> String {
if p.0.is_null() {
return String::new();
}
unsafe { p.to_string() }.unwrap_or_default()
}
fn instance_label(full: &str, service_type: &str, domain: &str) -> String {
let suffix = format!(".{service_type}.{domain}");
full.strip_suffix(&suffix).unwrap_or(full).to_string()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn instance_label_strips_type_and_domain_suffix() {
let label = instance_label("My Printer._ipp._tcp.local", "_ipp._tcp", "local");
assert_eq!(label, "My Printer");
}
#[test]
fn instance_label_handles_dotted_instance_name() {
let label = instance_label("host.example._http._tcp.local", "_http._tcp", "local");
assert_eq!(label, "host.example");
}
#[test]
fn instance_label_returns_full_when_suffix_absent() {
let full = "Something._other._tcp.local";
assert_eq!(instance_label(full, "_http._tcp", "local"), full);
}
}