use crate::proto; use crate::proto::utils::{RepeatedDeserialize, RepeatedMessageModel, RepeatedSerialize}; use anyhow::{Error, anyhow}; use hickory_net::xfer::Protocol; use hickory_resolver::config::{ConnectionConfig, NameServerConfig, ProtocolConfig}; use serde::de::IntoDeserializer; use serde::{Deserialize, Deserializer, de}; use serde_with::{DeserializeFromStr, SerializeDisplay}; use std::fmt::{Display, Formatter}; use std::net::{IpAddr, Ipv6Addr, SocketAddr}; use std::str::FromStr; use url::Url; #[derive(Debug, Copy, Clone, PartialEq, Eq, Hash, SerializeDisplay, DeserializeFromStr)] pub struct NameServerAddr { pub protocol: Protocol, pub addr: SocketAddr, } impl From for NameServerConfig { fn from(value: NameServerAddr) -> Self { let mut config = match value.protocol { Protocol::Udp => ConnectionConfig::udp(), Protocol::Tcp => ConnectionConfig::tcp(), _ => unimplemented!(), }; config.port = value.addr.port(); Self::new(value.addr.ip(), true, vec![config]) } } impl From<(IpAddr, &ConnectionConfig)> for NameServerAddr { fn from(value: (IpAddr, &ConnectionConfig)) -> Self { let (ip, config) = value; Self { protocol: config.protocol.to_protocol(), addr: SocketAddr::new(ip, config.port), } } } impl TryFrom<&Url> for NameServerAddr { type Error = Error; fn try_from(url: &Url) -> Result { let protocol = match Protocol::deserialize(url.scheme().into_deserializer()) .map_err(|e: de::value::Error| anyhow!("invalid protocol '{}': {}", url.scheme(), e))? { Protocol::Udp => ProtocolConfig::Udp, Protocol::Tcp => ProtocolConfig::Tcp, p => return Err(anyhow!("unsupported protocol: {}", p)), }; let host = url.host_str().ok_or(anyhow!("host not found"))?; let port = url.port().unwrap_or(protocol.default_port()); let addr = if let Ok(addr) = IpAddr::from_str(host) { SocketAddr::new(addr, port) } else { return Err(anyhow!("invalid address: {}", host)); }; Ok(Self { protocol: protocol.to_protocol(), addr, }) } } impl TryFrom<&proto::common::Url> for NameServerAddr { type Error = Error; fn try_from(value: &proto::common::Url) -> Result { (&Url::try_from(value)?).try_into() } } impl From<&NameServerAddr> for Url { fn from(value: &NameServerAddr) -> Self { Url::parse(&format!("{}://{}", value.protocol, value.addr)).unwrap() } } impl From<&NameServerAddr> for proto::common::Url { fn from(value: &NameServerAddr) -> Self { Url::from(value).into() } } impl From for Url { fn from(value: NameServerAddr) -> Self { (&value).into() } } impl From for proto::common::Url { fn from(value: NameServerAddr) -> Self { (&value).into() } } impl FromStr for NameServerAddr { type Err = Error; fn from_str(s: &str) -> Result { (&Url::parse(s)?).try_into() } } impl Display for NameServerAddr { fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { f.write_str(Url::from(*self).as_str()) } } pub type NameServerAddrGroup = RepeatedMessageModel; impl From<&NameServerConfig> for NameServerAddrGroup { fn from(value: &NameServerConfig) -> Self { value .connections .iter() .map(|c| (value.ip, c).into()) .collect() } } impl From for NameServerAddrGroup { fn from(value: SocketAddr) -> Self { vec![ NameServerAddr { protocol: Protocol::Udp, addr: value, }, NameServerAddr { protocol: Protocol::Tcp, addr: value, }, ] .into() } } impl From for NameServerAddrGroup { fn from(value: IpAddr) -> Self { SocketAddr::new(value, 53).into() } } impl From for NameServerAddrGroup { fn from(value: u16) -> Self { SocketAddr::new(Ipv6Addr::UNSPECIFIED.into(), value).into() } } impl RepeatedSerialize for NameServerAddr {} impl<'de> RepeatedDeserialize<'de> for NameServerAddr { fn deserialize(deserializer: D) -> Result where D: Deserializer<'de>, { #[derive(Deserialize)] #[serde(untagged)] enum Candidate { NameServerAddr(NameServerAddr), U16(u16), IpAddr(IpAddr), SocketAddr(SocketAddr), } let items = Vec::::deserialize(deserializer)?; let items = items .into_iter() .flat_map(|item| -> NameServerAddrGroup { match item { Candidate::NameServerAddr(addr) => vec![addr].into(), Candidate::U16(port) => port.into(), Candidate::IpAddr(ip) => ip.into(), Candidate::SocketAddr(addr) => addr.into(), } }) .collect(); Ok(items) } }