dyns 0.3.0

DNS discovery and resolver support for DHTTP applications
Documentation
use std::{
    net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr},
    num::{NonZero, NonZeroU32},
    pin::Pin,
    sync::{Arc, Weak},
    task::{Context, Poll},
    time::Duration,
};

use dashmap::DashMap;
use futures::{Stream, StreamExt};
use snafu::Snafu;
use socket2::{Domain, Socket, Type};
use tokio::{io, net::UdpSocket, task::JoinSet, time};

use super::if_nametoindex::if_nametoindex;
use crate::core::parser::{
    packet::{Packet, be_packet},
    record::endpoint::EndpointAddr,
};

#[derive(Debug)]
pub struct MdnsSocket {
    udp: UdpSocket,
    ip: IpAddr,
    nic: String,
}

const MULTICAST_PORT: u16 = 5353;
const MULTICAST_ADDR_V4: Ipv4Addr = Ipv4Addr::new(224, 0, 0, 251);
const MULTICAST_ADDR_V6: Ipv6Addr = Ipv6Addr::new(0xff02, 0, 0, 0, 0, 0, 0, 0xfb);
const MAX_DEQUE_SIZE: usize = 64;

impl MdnsSocket {
    pub fn new(device: &str, ip: IpAddr) -> io::Result<Self> {
        tracing::trace!(target: "mdns", device, %ip, "add mdns device");
        let socket = match ip {
            #[cfg_attr(
                not(any(target_os = "android", target_os = "fuchsia", target_os = "linux")),
                allow(clippy::unused_variables)
            )]
            IpAddr::V4(ip) => {
                let socket = Socket::new(Domain::IPV4, Type::DGRAM, None)?;
                socket.set_nonblocking(true)?;
                socket.set_reuse_address(true)?;
                #[cfg(not(target_os = "windows"))]
                socket.set_reuse_port(true)?;

                let bind = SocketAddr::new(IpAddr::V4(Ipv4Addr::UNSPECIFIED), MULTICAST_PORT);
                socket.bind(&bind.into())?;
                #[cfg(any(target_os = "android", target_os = "fuchsia", target_os = "linux"))]
                socket.bind_device(Some(device.as_bytes()))?;
                #[cfg(any(
                    target_os = "ios",
                    target_os = "visionos",
                    target_os = "macos",
                    target_os = "tvos",
                    target_os = "watchos",
                ))]
                {
                    let ifindex = NonZeroU32::new(if_nametoindex(device)?).ok_or_else(|| {
                        io::Error::new(
                            io::ErrorKind::InvalidInput,
                            "interface index must be non-zero",
                        )
                    })?;
                    socket.bind_device_by_index_v4(Some(ifindex))?;
                }
                // Always enable multicast loopback so that mDNS services on the
                // same host (but in different processes) can communicate.
                socket.set_multicast_loop_v4(true)?;
                // 使用接口自身的 IP 加入多播组,确保该 socket 只在对应接口上收发 mDNS 报文,
                // 避免以 UNSPECIFIED(0.0.0.0) 加入导致 lo0 等 socket 响应其他接口的查询。
                // FIXME:这个改动在树莓派上可能有问题,树莓派上只能 join 0.0.0.0
                socket.join_multicast_v4(&MULTICAST_ADDR_V4, &ip)?;
                socket.set_multicast_if_v4(&ip)?;
                socket
            }
            IpAddr::V6(_ip) => {
                let socket = Socket::new(Domain::IPV6, Type::DGRAM, None)?;
                socket.set_nonblocking(true)?;
                socket.set_reuse_address(true)?;
                #[cfg(not(target_os = "windows"))]
                socket.set_reuse_port(true)?;

                let bind = SocketAddr::new(IpAddr::V6(Ipv6Addr::UNSPECIFIED), MULTICAST_PORT);
                socket.bind(&bind.into())?;
                #[cfg(any(target_os = "android", target_os = "fuchsia", target_os = "linux"))]
                socket.bind_device(Some(device.as_bytes()))?;
                // Always enable multicast loopback so that mDNS services on the
                // same host (but in different processes) can communicate.
                socket.set_multicast_loop_v6(true)?;
                // TODO: 外面传进来
                let ifindex = NonZeroU32::new(if_nametoindex(device)?).ok_or_else(|| {
                    io::Error::new(
                        io::ErrorKind::InvalidInput,
                        "interface index must be non-zero",
                    )
                })?;
                #[cfg(any(
                    target_os = "ios",
                    target_os = "visionos",
                    target_os = "macos",
                    target_os = "tvos",
                    target_os = "watchos",
                ))]
                socket.bind_device_by_index_v6(Some(ifindex))?;
                socket.join_multicast_v6(&MULTICAST_ADDR_V6, ifindex.get())?;
                socket.set_multicast_if_v6(ifindex.get())?;

                socket
            }
        };

        Ok(Self {
            udp: UdpSocket::from_std(socket.into())?,
            ip,
            nic: device.to_string(),
        })
    }

    pub async fn receive(&self) -> io::Result<(SocketAddr, Packet)> {
        loop {
            let mut recv_buffer = [0u8; 2048];

            let (size, source) = self.udp.recv_from(&mut recv_buffer).await?;

            let Ok((_remain, packet)) = be_packet(&recv_buffer[..size]) else {
                continue;
            };

            return Ok((source, packet));
        }
    }

    pub async fn broadcast_packet(&self, packet: Packet) -> io::Result<()> {
        let buf = packet.to_bytes();
        let target: SocketAddr = match self.udp.local_addr()?.ip() {
            IpAddr::V4(_) => (MULTICAST_ADDR_V4, MULTICAST_PORT).into(),
            IpAddr::V6(_) => (MULTICAST_ADDR_V6, MULTICAST_PORT).into(),
        };
        self.udp.send_to(&buf, target).await?;

        Ok(())
    }
}

#[allow(clippy::type_complexity)]
pub struct PacketRouter {
    requests: (
        flume::Sender<(SocketAddr, Packet)>,
        flume::Receiver<(SocketAddr, Packet)>,
    ),
    responses: (
        flume::Sender<(SocketAddr, Packet)>,
        flume::Receiver<(SocketAddr, Packet)>,
    ),
    queries: DashMap<NonZero<u16>, flume::Sender<(SocketAddr, Packet)>>,
}

impl PacketRouter {
    pub fn new() -> Self {
        Self {
            requests: flume::bounded(MAX_DEQUE_SIZE),
            responses: flume::bounded(MAX_DEQUE_SIZE),
            queries: DashMap::new(),
        }
    }

    pub fn receive_query(
        &self,
    ) -> impl Future<Output = Result<(SocketAddr, Packet), flume::RecvError>> + Send + use<> {
        self.requests.1.clone().into_recv_async()
    }

    pub fn receive_boardcast(
        &self,
    ) -> impl Future<Output = Result<(SocketAddr, Packet), flume::RecvError>> + Send + use<> {
        self.responses.1.clone().into_recv_async()
    }

    pub async fn register_query(
        self: &Arc<Self>,
        query_id: NonZero<u16>,
    ) -> impl Stream<Item = (SocketAddr, Packet)> {
        struct Responses {
            query_id: NonZero<u16>,
            router: Weak<PacketRouter>,
            recver: flume::r#async::RecvStream<'static, (SocketAddr, Packet)>,
        }

        impl Stream for Responses {
            type Item = (SocketAddr, Packet);

            fn poll_next(
                mut self: Pin<&mut Self>,
                cx: &mut Context<'_>,
            ) -> Poll<Option<Self::Item>> {
                Pin::new(&mut self.recver).poll_next(cx)
            }
        }

        impl Drop for Responses {
            fn drop(&mut self) {
                if let Some(router) = self.router.upgrade() {
                    router.queries.remove(&self.query_id);
                }
            }
        }

        let (tx, rx) = flume::bounded(MAX_DEQUE_SIZE);
        self.queries.insert(query_id, tx);

        Responses {
            query_id,
            router: Arc::downgrade(self),
            recver: rx.into_stream(),
        }
    }

    pub fn deliver(&self, source: SocketAddr, packet: Packet) {
        match (packet.is_query(), packet.id()) {
            (true, 0) => {
                if self.responses.0.try_send((source, packet.clone())).is_err() {
                    // Queue is full, remove oldest message (FIFO)
                    let _ = self.responses.1.try_recv();
                    // Try to send again after removing oldest
                    let _ = self.responses.0.try_send((source, packet));
                }
            }
            (true, query_id) => match self.queries.get(&NonZero::new(query_id).unwrap()) {
                Some(tx) => {
                    if let Err(error) = tx.try_send((source, packet)) {
                        tracing::debug!(
                            target: "mdns",
                            %query_id, %error,
                            "failed to route response for query id"
                        );
                    }
                }
                None => tracing::debug!(
                    target: "mdns",
                    %query_id,
                    "received response for query id, but no such query registered"
                ),
            },
            (false, _) => {
                if self.requests.0.try_send((source, packet.clone())).is_err() {
                    // Queue is full, remove oldest message (FIFO)
                    let _ = self.requests.1.try_recv();
                    // Try to send again after removing oldest
                    let _ = self.requests.0.try_send((source, packet));
                }
            }
        }
    }
}

#[derive(Debug)]
pub struct MdnsProtocol {
    socket: Arc<MdnsSocket>,
    router: Weak<PacketRouter>,
}

#[derive(Debug, Snafu)]
#[snafu(display("mDNS socket is not listening"))]
pub struct Disconnected;

impl From<Disconnected> for io::Error {
    fn from(error: Disconnected) -> Self {
        io::Error::new(io::ErrorKind::NotConnected, error)
    }
}

impl MdnsProtocol {
    pub fn new(
        device: &str,
        ip: IpAddr,
    ) -> io::Result<(Self, impl Future<Output = ()> + Send + use<>)> {
        let socket = Arc::new(MdnsSocket::new(device, ip)?);
        let router = Arc::new(PacketRouter::new());

        let route = {
            let socket = socket.clone();
            let router = router.clone();
            async move {
                while let Ok((source, packet)) = socket.receive().await {
                    router.deliver(source, packet);
                }
            }
        };
        let protocol = Self {
            socket,
            router: Arc::downgrade(&router),
        };
        Ok((protocol, route))
    }

    pub async fn broadcast_packet(&self, packet: Packet) -> io::Result<()> {
        self.socket.broadcast_packet(packet).await
    }

    pub fn bound_ip(&self) -> IpAddr {
        self.socket.ip
    }

    pub fn bound_nic(&self) -> &str {
        &self.socket.nic
    }

    pub async fn query(
        self: &Arc<Self>,
        local_name: String,
    ) -> io::Result<(SocketAddr, Vec<EndpointAddr>)> {
        let router = self.router.upgrade().ok_or(Disconnected)?;

        let packet = Packet::query_with_id(local_name.clone());
        let query_id = NonZero::new(packet.id()).ok_or_else(|| {
            io::Error::new(io::ErrorKind::InvalidInput, "Query id should not be 0")
        })?;

        let mut packets = router.register_query(query_id).await;
        let mut broadcast_tasks = JoinSet::new();

        for _ in 0..3 {
            _ = broadcast_tasks.spawn({
                let this = self.clone();
                let packet = packet.clone();
                async move { this.broadcast_packet(packet).await }
            });

            if let Ok(Some((source, packet))) =
                time::timeout(Duration::from_millis(300), packets.next()).await
            {
                use crate::core::parser::record::RData::*;
                let endpoints = packet
                    .answers
                    .iter()
                    .inspect(|answer| {
                        tracing::debug!(target: "mdns", ?answer, "recv response");
                    })
                    .filter(|answer| {
                        if answer.name() != local_name {
                            tracing::debug!(
                                target: "mdns",
                                answer_name = answer.name(),
                                local_name,
                                "ignored answer for different service name",
                            );
                        }
                        answer.name() == local_name
                    })
                    .filter_map(|answer| match answer.data() {
                        E(e) => Some(e.clone()),
                        _ => {
                            tracing::debug!(target: "mdns", ?answer, "ignored record");
                            None
                        }
                    })
                    .collect::<Vec<_>>();

                if !endpoints.is_empty() {
                    return Ok((source, endpoints));
                }
            }
        }

        broadcast_tasks.abort_all();
        Err(io::ErrorKind::TimedOut.into())
    }

    pub async fn receive_query(&self) -> Result<(SocketAddr, Packet), Disconnected> {
        let router = self.router.upgrade().ok_or(Disconnected)?;

        router.receive_query().await.map_err(|_| Disconnected)
    }
    pub async fn receive_boardcast(&self) -> Result<(SocketAddr, Packet), Disconnected> {
        let router = self.router.upgrade().ok_or(Disconnected)?;

        router.receive_boardcast().await.map_err(|_| Disconnected)
    }
}