diff --git a/easytier/src/dns/client_mgr.rs b/easytier/src/dns/client_mgr.rs index 4a99d55c..ebb354af 100644 --- a/easytier/src/dns/client_mgr.rs +++ b/easytier/src/dns/client_mgr.rs @@ -1,4 +1,5 @@ -use crate::dns::utils::{DirtyFlag, DirtyState, NameServerAddr}; +use crate::dns::utils::addr::NameServerAddr; +use crate::dns::utils::dirty::{DirtyFlag, DirtyState}; use crate::dns::zone::{Zone, ZoneGroup}; use crate::proto::dns::DnsClientMgrRpc; use crate::proto::dns::{DnsSnapshot, HeartbeatRequest, HeartbeatResponse}; diff --git a/easytier/src/dns/config/dns.rs b/easytier/src/dns/config/dns.rs index 2af4c438..8288b731 100644 --- a/easytier/src/dns/config/dns.rs +++ b/easytier/src/dns/config/dns.rs @@ -2,7 +2,8 @@ use crate::common::global_ctx::GlobalCtx; use crate::dns::config::policy::DnsPolicyConfig; use crate::dns::config::zone::ZoneConfig; use crate::dns::config::{DNS_DEFAULT_ADDRESS, DNS_DEFAULT_TLD}; -use crate::dns::utils::{parse, NameServerAddrGroup}; +use crate::dns::utils::addr::NameServerAddrGroup; +use crate::dns::utils::parse; use crate::proto::dns::GetExportConfigResponse; use derivative::Derivative; use gethostname::gethostname; diff --git a/easytier/src/dns/config/mod.rs b/easytier/src/dns/config/mod.rs index b2e68789..86331ae4 100644 --- a/easytier/src/dns/config/mod.rs +++ b/easytier/src/dns/config/mod.rs @@ -1,4 +1,4 @@ -use crate::dns::utils::NameServerAddr; +use crate::dns::utils::addr::NameServerAddr; use hickory_proto::rr::LowerName; use hickory_proto::xfer::Protocol; use std::net::{Ipv4Addr, SocketAddr, SocketAddrV4}; diff --git a/easytier/src/dns/config/zone.rs b/easytier/src/dns/config/zone.rs index 1ea3e013..2bfe45e5 100644 --- a/easytier/src/dns/config/zone.rs +++ b/easytier/src/dns/config/zone.rs @@ -1,5 +1,5 @@ use crate::dns::config::policy::{DnsExportPolicy, ZonePolicyConfig}; -use crate::dns::utils::NameServerAddrGroup; +use crate::dns::utils::addr::NameServerAddrGroup; use crate::dns::zone::Zone; use crate::proto::dns::ZoneData; use derivative::Derivative; diff --git a/easytier/src/dns/mod.rs b/easytier/src/dns/mod.rs index b7451832..3fc5bb29 100644 --- a/easytier/src/dns/mod.rs +++ b/easytier/src/dns/mod.rs @@ -3,5 +3,5 @@ mod client_mgr; pub mod config; mod peer_mgr; mod server; -pub mod utils; +mod utils; pub mod zone; diff --git a/easytier/src/dns/peer_mgr.rs b/easytier/src/dns/peer_mgr.rs index 2bb3b6eb..b04cfe77 100644 --- a/easytier/src/dns/peer_mgr.rs +++ b/easytier/src/dns/peer_mgr.rs @@ -1,7 +1,7 @@ use crate::common::config::ConfigLoader; use crate::common::PeerId; use crate::dns::config::{DnsExportConfig, DnsGlobalCtxExt}; -use crate::dns::utils::{DirtyFlag, DirtyState}; +use crate::dns::utils::dirty::{DirtyFlag, DirtyState}; use crate::dns::zone::ZoneGroup; use crate::peer_center::instance::PeerCenterPeerManagerTrait; use crate::peers::peer_manager::PeerManager; diff --git a/easytier/src/dns/server.rs b/easytier/src/dns/server.rs index 4f2005b1..a414f910 100644 --- a/easytier/src/dns/server.rs +++ b/easytier/src/dns/server.rs @@ -1,5 +1,7 @@ -use super::utils::NameServerAddr; use crate::dns::client_mgr::DnsClientMgr; +use crate::dns::utils::addr::NameServerAddr; +use crate::peers::peer_manager::PeerManager; +use crate::proto::dns::DnsClientMgrRpcServer; use derivative::Derivative; use derive_more::{Deref, DerefMut, From, Into}; use hickory_proto::rr::Record; @@ -18,8 +20,6 @@ use std::{sync::Arc, time::Duration}; use tokio::net::{TcpListener, UdpSocket}; use tokio::{sync::RwLock, task::JoinHandle}; use tokio_util::sync::CancellationToken; -use crate::peers::peer_manager::PeerManager; -use crate::proto::dns::DnsClientMgrRpcServer; #[derive(Clone)] pub struct DynamicCatalog { diff --git a/easytier/src/dns/utils.rs b/easytier/src/dns/utils.rs deleted file mode 100644 index 2eca4187..00000000 --- a/easytier/src/dns/utils.rs +++ /dev/null @@ -1,309 +0,0 @@ -use crate::dns::config::DNS_SUPPORTED_PROTOCOLS; -use crate::proto; -use crate::proto::utils::RepeatedMessageModel; -use anyhow::{anyhow, Error}; -use derive_more::{Deref, DerefMut}; -use hickory_proto::rr::{LowerName, RecordType}; -use hickory_proto::xfer::Protocol; -use hickory_resolver::config::{NameServerConfig, NameServerConfigGroup}; -use hickory_server::authority::{ - Authority, LookupControlFlow, LookupObject, LookupOptions, MessageRequest, UpdateResult, - ZoneType, -}; -use hickory_server::server::RequestInfo; -use idna::AsciiDenyList; -use itertools::Itertools; -use serde_with::{DeserializeFromStr, SerializeDisplay}; -use std::fmt::{Display, Formatter}; -use std::net::{IpAddr, SocketAddr}; -use std::str::FromStr; -use std::sync::atomic::{AtomicBool, Ordering}; -use tokio::sync::Notify; -use url::Url; - -pub fn sanitize(name: &str) -> String { - let dot = name.ends_with('.'); - let mut name = idna::domain_to_ascii_cow(name.as_ref(), AsciiDenyList::EMPTY) - .unwrap_or_default() - .into_owned() - .to_lowercase() - .split('.') - .map(|label| { - label - .chars() - .map(|c| if c.is_ascii_alphanumeric() { c } else { '-' }) - .take(63) - .collect::() - .trim_matches('-') - .to_string() - }) - .filter(|label| !label.is_empty()) - .collect_vec() - .join("."); - name.truncate(253); - if dot { - name.push('.'); - } - name -} - -pub fn parse(name: &str) -> LowerName { - if let Ok(name) = name.parse() { - name - } else { - let sanitized = sanitize(name); - tracing::debug!("invalid hostname: {}, sanitized to: {}", name, sanitized); - sanitized.parse().unwrap_or_default() - } -} - -#[derive(Debug, Copy, Clone, PartialEq, Eq, Hash, SerializeDisplay, DeserializeFromStr)] -pub struct NameServerAddr { - pub(super) protocol: Protocol, - pub(super) addr: SocketAddr, -} - -impl From for NameServerConfig { - fn from(value: NameServerAddr) -> Self { - Self::new(value.addr, value.protocol) - } -} - -impl From for NameServerAddr { - fn from(value: NameServerConfig) -> Self { - Self { - protocol: value.protocol, - addr: value.socket_addr, - } - } -} - -impl From for NameServerAddr { - fn from(value: SocketAddr) -> Self { - Self { - protocol: Protocol::Udp, - addr: value, - } - } -} - -impl From for NameServerAddr { - fn from(value: IpAddr) -> Self { - SocketAddr::new(value, 53).into() - } -} - -impl From for Url { - fn from(value: NameServerAddr) -> Self { - Url::parse(&format!("{}://{}", value.protocol, value.addr)).unwrap() - } -} - -impl TryFrom<&Url> for NameServerAddr { - type Error = Error; - - fn try_from(value: &Url) -> Result { - let scheme = value.scheme(); - let protocol = *DNS_SUPPORTED_PROTOCOLS - .iter() - .find(|p| p.to_string() == scheme) - .ok_or(anyhow!("unsupported scheme: {}", scheme))?; - let addr = value.host_str().ok_or(anyhow!("host not found"))?; - let addr = addr - .trim_start_matches('[') - .trim_end_matches(']') - .parse::() - .map_err(|e| anyhow!("invalid ip address '{}': {}", addr, e))?; - let port = if let Some(port) = value.port() { - port - } else { - match protocol { - Protocol::Udp | Protocol::Tcp => 53, - _ => return Err(anyhow!("port not found")), - } - }; - - Ok(Self { - protocol, - addr: SocketAddr::new(addr, port), - }) - } -} - -impl From for proto::common::Url { - fn from(value: NameServerAddr) -> Self { - Url::from(value).into() - } -} - -impl TryFrom<&proto::common::Url> for NameServerAddr { - type Error = Error; - - fn try_from(value: &proto::common::Url) -> Result { - Self::try_from(&Url::try_from(value)?) - } -} - -impl FromStr for NameServerAddr { - type Err = Error; - fn from_str(s: &str) -> Result { - macro_rules! try_parse { - ($($t:ty),+) => { - $( if let Ok(v) = s.parse::<$t>() { return Ok(v.into()); } )+ - }; - } - - try_parse!(IpAddr, SocketAddr); - - (&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(super) type NameServerAddrGroup = RepeatedMessageModel; - -impl From for NameServerConfigGroup { - fn from(value: NameServerAddrGroup) -> Self { - value.into_iter().map_into().collect_vec().into() - } -} - -impl From for NameServerAddrGroup { - fn from(value: NameServerConfigGroup) -> Self { - value - .into_inner() - .into_iter() - .map_into() - .collect_vec() - .into() - } -} - -#[derive(Deref, DerefMut)] -pub struct ChainedAuthority(pub(super) A) -where - A: Authority, - A::Lookup: LookupObject + 'static; - -impl From for ChainedAuthority -where - A: Authority, - A::Lookup: LookupObject + 'static, -{ - fn from(value: A) -> Self { - Self(value) - } -} - -#[async_trait::async_trait] -impl Authority for ChainedAuthority -where - A: Authority, - A::Lookup: LookupObject + 'static, -{ - type Lookup = A::Lookup; - - #[inline] - fn zone_type(&self) -> ZoneType { - self.0.zone_type() - } - #[inline] - fn is_axfr_allowed(&self) -> bool { - self.0.is_axfr_allowed() - } - #[inline] - async fn update(&self, update: &MessageRequest) -> UpdateResult { - self.0.update(update).await - } - #[inline] - fn origin(&self) -> &LowerName { - self.0.origin() - } - #[inline] - async fn lookup( - &self, - name: &LowerName, - rtype: RecordType, - lookup_options: LookupOptions, - ) -> LookupControlFlow { - self.0.lookup(name, rtype, lookup_options).await - } - #[inline] - async fn consult( - &self, - name: &LowerName, - rtype: RecordType, - lookup_options: LookupOptions, - last_result: LookupControlFlow>, - ) -> LookupControlFlow> { - if let Some(Ok(l)) = last_result.map_result() { - LookupControlFlow::Break(Ok(l)) - } else { - self.0 - .lookup(name, rtype, lookup_options) - .await - .map(|l| Box::new(l) as _) - } - } - #[inline] - async fn search( - &self, - request_info: RequestInfo<'_>, - lookup_options: LookupOptions, - ) -> LookupControlFlow { - self.0.search(request_info, lookup_options).await - } - #[inline] - async fn get_nsec_records( - &self, - name: &LowerName, - lookup_options: LookupOptions, - ) -> LookupControlFlow { - self.0.get_nsec_records(name, lookup_options).await - } -} - -#[derive(Debug, Deref, DerefMut)] -pub(super) struct DirtyState { - #[deref] - #[deref_mut] - flags: T, - pub notify: Notify, -} - -impl Default for DirtyState { - fn default() -> Self { - Self { - flags: T::default(), - notify: Notify::new(), - } - } -} - -#[derive(Debug)] -pub(super) struct DirtyFlag(AtomicBool); - -impl DirtyFlag { - pub fn new(value: bool) -> Self { - Self(AtomicBool::new(value)) - } - - pub fn mark(&self) { - self.0.store(true, Ordering::Release); - } - - pub fn reset(&self) -> bool { - self.0.swap(false, Ordering::Acquire) - } -} - -impl Default for DirtyFlag { - fn default() -> Self { - Self::new(true) - } -} diff --git a/easytier/src/dns/utils/addr.rs b/easytier/src/dns/utils/addr.rs new file mode 100644 index 00000000..055c14ff --- /dev/null +++ b/easytier/src/dns/utils/addr.rs @@ -0,0 +1,139 @@ +use crate::dns::config::DNS_SUPPORTED_PROTOCOLS; +use crate::proto; +use crate::proto::utils::RepeatedMessageModel; +use anyhow::{anyhow, Error}; +use hickory_proto::xfer::Protocol; +use hickory_resolver::config::{NameServerConfig, NameServerConfigGroup}; +use itertools::Itertools; +use serde_with::{DeserializeFromStr, SerializeDisplay}; +use std::fmt::{Display, Formatter}; +use std::net::{IpAddr, 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 { + Self::new(value.addr, value.protocol) + } +} + +impl From for NameServerAddr { + fn from(value: NameServerConfig) -> Self { + Self { + protocol: value.protocol, + addr: value.socket_addr, + } + } +} + +impl From for NameServerAddr { + fn from(value: SocketAddr) -> Self { + Self { + protocol: Protocol::Udp, + addr: value, + } + } +} + +impl From for NameServerAddr { + fn from(value: IpAddr) -> Self { + SocketAddr::new(value, 53).into() + } +} + +impl From for Url { + fn from(value: NameServerAddr) -> Self { + Url::parse(&format!("{}://{}", value.protocol, value.addr)).unwrap() + } +} + +impl TryFrom<&Url> for NameServerAddr { + type Error = Error; + + fn try_from(value: &Url) -> Result { + let scheme = value.scheme(); + let protocol = *DNS_SUPPORTED_PROTOCOLS + .iter() + .find(|p| p.to_string() == scheme) + .ok_or(anyhow!("unsupported scheme: {}", scheme))?; + let addr = value.host_str().ok_or(anyhow!("host not found"))?; + let addr = addr + .trim_start_matches('[') + .trim_end_matches(']') + .parse::() + .map_err(|e| anyhow!("invalid ip address '{}': {}", addr, e))?; + let port = if let Some(port) = value.port() { + port + } else { + match protocol { + Protocol::Udp | Protocol::Tcp => 53, + _ => return Err(anyhow!("port not found")), + } + }; + + Ok(Self { + protocol, + addr: SocketAddr::new(addr, port), + }) + } +} + +impl From for proto::common::Url { + fn from(value: NameServerAddr) -> Self { + Url::from(value).into() + } +} + +impl TryFrom<&proto::common::Url> for NameServerAddr { + type Error = Error; + + fn try_from(value: &proto::common::Url) -> Result { + Self::try_from(&Url::try_from(value)?) + } +} + +impl FromStr for NameServerAddr { + type Err = Error; + fn from_str(s: &str) -> Result { + macro_rules! try_parse { + ($($t:ty),+) => { + $( if let Ok(v) = s.parse::<$t>() { return Ok(v.into()); } )+ + }; + } + + try_parse!(IpAddr, SocketAddr); + + (&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 for NameServerConfigGroup { + fn from(value: NameServerAddrGroup) -> Self { + value.into_iter().map_into().collect_vec().into() + } +} + +impl From for NameServerAddrGroup { + fn from(value: NameServerConfigGroup) -> Self { + value + .into_inner() + .into_iter() + .map_into() + .collect_vec() + .into() + } +} diff --git a/easytier/src/dns/utils/authority.rs b/easytier/src/dns/utils/authority.rs new file mode 100644 index 00000000..c2f5aa57 --- /dev/null +++ b/easytier/src/dns/utils/authority.rs @@ -0,0 +1,81 @@ +use derive_more::{Deref, DerefMut, From}; +use hickory_proto::rr::{LowerName, RecordType}; +use hickory_server::authority::{ + Authority, LookupControlFlow, LookupObject, LookupOptions, MessageRequest, UpdateResult, + ZoneType, +}; +use hickory_server::server::RequestInfo; + +#[derive(From, Deref, DerefMut)] +pub struct ChainedAuthority(A) +where + A: Authority, + A::Lookup: LookupObject + 'static; + +#[async_trait::async_trait] +impl Authority for ChainedAuthority +where + A: Authority, + A::Lookup: LookupObject + 'static, +{ + type Lookup = A::Lookup; + + #[inline] + fn zone_type(&self) -> ZoneType { + self.0.zone_type() + } + #[inline] + fn is_axfr_allowed(&self) -> bool { + self.0.is_axfr_allowed() + } + #[inline] + async fn update(&self, update: &MessageRequest) -> UpdateResult { + self.0.update(update).await + } + #[inline] + fn origin(&self) -> &LowerName { + self.0.origin() + } + #[inline] + async fn lookup( + &self, + name: &LowerName, + rtype: RecordType, + lookup_options: LookupOptions, + ) -> LookupControlFlow { + self.0.lookup(name, rtype, lookup_options).await + } + #[inline] + async fn consult( + &self, + name: &LowerName, + rtype: RecordType, + lookup_options: LookupOptions, + last_result: LookupControlFlow>, + ) -> LookupControlFlow> { + if let Some(Ok(l)) = last_result.map_result() { + LookupControlFlow::Break(Ok(l)) + } else { + self.0 + .lookup(name, rtype, lookup_options) + .await + .map(|l| Box::new(l) as _) + } + } + #[inline] + async fn search( + &self, + request_info: RequestInfo<'_>, + lookup_options: LookupOptions, + ) -> LookupControlFlow { + self.0.search(request_info, lookup_options).await + } + #[inline] + async fn get_nsec_records( + &self, + name: &LowerName, + lookup_options: LookupOptions, + ) -> LookupControlFlow { + self.0.get_nsec_records(name, lookup_options).await + } +} diff --git a/easytier/src/dns/utils/dirty.rs b/easytier/src/dns/utils/dirty.rs new file mode 100644 index 00000000..3faa3247 --- /dev/null +++ b/easytier/src/dns/utils/dirty.rs @@ -0,0 +1,30 @@ +use derivative::Derivative; +use derive_more::{Deref, DerefMut}; +use std::sync::atomic::{AtomicBool, Ordering}; +use tokio::sync::Notify; + +#[derive(Debug, Default, Deref, DerefMut)] +pub struct DirtyState { + #[deref] + #[deref_mut] + flags: T, + pub notify: Notify, +} + +#[derive(Derivative, Debug)] +#[derivative(Default)] +pub struct DirtyFlag(#[derivative(Default(value = "AtomicBool::new(true)"))] AtomicBool); + +impl DirtyFlag { + pub fn new(value: bool) -> Self { + Self(AtomicBool::new(value)) + } + + pub fn mark(&self) { + self.0.store(true, Ordering::Release); + } + + pub fn reset(&self) -> bool { + self.0.swap(false, Ordering::Acquire) + } +} diff --git a/easytier/src/dns/utils/mod.rs b/easytier/src/dns/utils/mod.rs new file mode 100644 index 00000000..4d54b157 --- /dev/null +++ b/easytier/src/dns/utils/mod.rs @@ -0,0 +1,43 @@ +use hickory_proto::rr::LowerName; +use idna::AsciiDenyList; +use itertools::Itertools; + +pub mod addr; +pub mod authority; +pub mod dirty; + +pub fn sanitize(name: &str) -> String { + let dot = name.ends_with('.'); + let mut name = idna::domain_to_ascii_cow(name.as_ref(), AsciiDenyList::EMPTY) + .unwrap_or_default() + .into_owned() + .to_lowercase() + .split('.') + .map(|label| { + label + .chars() + .map(|c| if c.is_ascii_alphanumeric() { c } else { '-' }) + .take(63) + .collect::() + .trim_matches('-') + .to_string() + }) + .filter(|label| !label.is_empty()) + .collect_vec() + .join("."); + name.truncate(253); + if dot { + name.push('.'); + } + name +} + +pub fn parse(name: &str) -> LowerName { + if let Ok(name) = name.parse() { + name + } else { + let sanitized = sanitize(name); + tracing::debug!("invalid hostname: {}, sanitized to: {}", name, sanitized); + sanitized.parse().unwrap_or_default() + } +} diff --git a/easytier/src/dns/zone.rs b/easytier/src/dns/zone.rs index 7dcf365c..4eafa181 100644 --- a/easytier/src/dns/zone.rs +++ b/easytier/src/dns/zone.rs @@ -1,5 +1,5 @@ use crate::common::dns::get_default_resolver_config; -use crate::dns::utils::NameServerAddr; +use crate::dns::utils::addr::NameServerAddr; use crate::proto; use crate::proto::utils::RepeatedMessageModel; use crate::utils::MapTryInto;