diff --git a/easytier/Cargo.toml b/easytier/Cargo.toml index bb1b0e75..33327cf9 100644 --- a/easytier/Cargo.toml +++ b/easytier/Cargo.toml @@ -50,6 +50,8 @@ time = "0.3" toml = "0.8.12" chrono = { version = "0.4.37", features = ["serde"] } +optional_struct = "0.5.2" + guarden = "0.1" delegate = "0.13.5" diff --git a/easytier/src/common/config.rs b/easytier/src/common/config.rs index ede1f525..75db4458 100644 --- a/easytier/src/common/config.rs +++ b/easytier/src/common/config.rs @@ -1,18 +1,3 @@ -use std::{ - hash::Hasher, - net::{IpAddr, SocketAddr}, - path::PathBuf, - sync::{Arc, Mutex}, -}; - -use anyhow::Context; -use base64::{Engine as _, prelude::BASE64_STANDARD}; -use clap::ValueEnum; -use clap::builder::PossibleValue; -use serde::{Deserialize, Serialize}; -use strum::{Display, EnumString, VariantArray}; -use tokio::io::AsyncReadExt as _; - use super::env_parser; use crate::utils::dns::sanitize; use crate::{ @@ -25,6 +10,89 @@ use crate::{ tunnel::{IpScheme, TunnelScheme, generate_digest_from_str}, utils, }; +use anyhow::Context; +use base64::{Engine as _, prelude::BASE64_STANDARD}; +use clap::ValueEnum; +use clap::builder::PossibleValue; +use derivative::Derivative; +use derive_more::{Constructor, Deref}; +use optional_struct::Applicable; +use serde::{Deserialize, Serialize}; +use std::fmt::{Debug, Display}; +use std::{ + hash::Hasher, + net::{IpAddr, SocketAddr}, + path::PathBuf, + sync::{Arc, Mutex}, +}; +use strum::{Display, EnumString, VariantArray}; +use tokio::io::AsyncReadExt as _; + +#[derive(Derivative, Debug, Clone, Constructor, Deref, Deserialize)] +#[derivative(PartialEq(bound = "Parsed: PartialEq"))] +#[serde(try_from = "Raw")] +#[serde( + bound = "Raw: Deserialize<'de>, as TryFrom>::Error: Display" +)] +pub struct ConfigBase +where + Raw: Applicable, + ConfigBase: TryFrom, +{ + #[derivative(PartialEq = "ignore")] + raw: Raw, + #[deref] + parsed: Parsed, + #[derivative(PartialEq = "ignore")] + data: Data, +} + +impl Serialize for ConfigBase +where + Raw: Applicable + Serialize, + ConfigBase: TryFrom, +{ + fn serialize(&self, serializer: S) -> Result + where + S: serde::Serializer, + { + self.raw.serialize(serializer) + } +} + +impl Default for ConfigBase +where + Raw: Applicable + Default, + ConfigBase: TryFrom, +{ + fn default() -> Self { + Raw::default().try_into().unwrap() + } +} + +impl ConfigBase +where + Raw: Applicable, + ConfigBase: TryFrom, +{ + pub fn into_raw(self) -> Raw { + self.raw + } + + pub fn into_parsed(self) -> Parsed { + self.parsed + } + + pub fn into_data(self) -> Data { + self.data + } + + pub fn update(self, config: Raw) -> Result>::Error> { + let mut raw = self.into_raw(); + config.apply_to_opt(&mut raw); + raw.try_into() + } +} pub type Flags = crate::proto::common::FlagsInConfig; @@ -554,7 +622,8 @@ struct Config { proxy_network: Option>, #[cfg(feature = "magic-dns")] - dns: Option, + #[serde(default)] + dns: DnsConfig, vpn_portal_config: Option, @@ -655,9 +724,9 @@ impl DnsConfigLoaderExt for TomlConfigLoader { cfg_select! { feature = "magic-dns" => { fn get_dns(&self) -> DnsConfig { - self.config.lock().unwrap().dns.clone().unwrap_or_default() + self.config.lock().unwrap().dns.clone() } - fn set_dns(&self, config: Option) { + fn set_dns(&self, config: DnsConfig) { self.config.lock().unwrap().dns = config; } } diff --git a/easytier/src/common/global_ctx.rs b/easytier/src/common/global_ctx.rs index e64673cc..4bc1e700 100644 --- a/easytier/src/common/global_ctx.rs +++ b/easytier/src/common/global_ctx.rs @@ -21,7 +21,7 @@ use super::{ stun::{StunInfoCollector, StunInfoCollectorTrait}, }; #[cfg(feature = "magic-dns")] -use crate::dns::config::{DnsConfigLoaderExt, DnsExportConfig, DnsGlobalCtxExt}; +use crate::dns::config::{DnsConfigLoaderExt, DnsExportConfig, DnsGlobalCtxExt, zone::ZoneConfig}; use crate::{ common::{ config::ProxyNetworkConfig, shrink_dashmap, stats_manager::StatsManager, @@ -461,7 +461,7 @@ impl GlobalCtx { } pub fn get_hostname(&self) -> String { - return self.hostname.lock().unwrap().clone(); + self.hostname.lock().unwrap().clone() } pub fn set_hostname(&self, hostname: String) { @@ -713,7 +713,7 @@ impl GlobalCtx { #[cfg(feature = "magic-dns")] impl DnsGlobalCtxExt for GlobalCtx { - fn dns_self_zone(&self) -> crate::dns::config::zone::ZoneConfig { + fn dns_self_zone(&self) -> ZoneConfig { let dns = self.config.get_dns(); let mut hostname = dns.name.to_string(); if hostname.is_empty() { @@ -728,7 +728,7 @@ impl DnsGlobalCtxExt for GlobalCtx { let ipv6 = self.get_ipv6().map(|ip| ip.address()); let ipv6 = ipv6.map(|a| vec![a]).unwrap_or_default(); - crate::dns::config::zone::ZoneConfig::dedicated(fqdn, ipv4, ipv6).unwrap() + ZoneConfig::dedicated(fqdn, ipv4, ipv6) } fn dns_export_config(&self) -> DnsExportConfig { @@ -736,13 +736,13 @@ impl DnsGlobalCtxExt for GlobalCtx { zones: self .dns_iter_zones() .filter(|z| z.policy.export.as_ref().is_some_and(|f| !f.disabled)) // TODO: check policies of parent zones - .map(Into::into) + .map(ZoneConfig::into_data) .collect(), } } - fn dns_iter_zones(&self) -> impl Iterator { - iter::once(self.dns_self_zone()).chain(self.config.get_dns().zones) + fn dns_iter_zones(&self) -> impl Iterator { + iter::once(self.dns_self_zone()).chain(self.config.get_dns().into_parsed().zones) } } diff --git a/easytier/src/dns/config/dns.rs b/easytier/src/dns/config/dns.rs index 5b9333fc..811238b6 100644 --- a/easytier/src/dns/config/dns.rs +++ b/easytier/src/dns/config/dns.rs @@ -1,33 +1,48 @@ +use crate::common::config::ConfigBase; use crate::dns::config::policy::DnsPolicyConfig; use crate::dns::config::zone::ZoneConfig; use crate::dns::config::{DNS_DEFAULT_ADDRESSES, DNS_DEFAULT_DOMAIN}; use crate::dns::utils::addr::NameServerAddrGroup; use crate::proto::dns::GetExportConfigResponse; -use derivative::Derivative; use hickory_proto::rr::LowerName; +use optional_struct::{Applicable, optional_struct}; use serde::{Deserialize, Serialize}; use std::collections::HashMap; -#[derive(Derivative, Debug, Clone, Deserialize, Serialize, PartialEq)] -#[derivative(Default)] -#[serde(default)] -pub struct DnsConfig { +#[optional_struct(DnsConfigRaw)] +#[derive(Debug, Clone, Default, PartialEq, Deserialize, Serialize)] +pub struct DnsConfigParsed { #[serde(rename = "zone")] pub zones: Vec, + #[optional_skip_wrap] #[serde(flatten)] pub policies: HashMap, pub name: LowerName, - #[derivative(Default(value = "DNS_DEFAULT_DOMAIN.clone()"))] pub domain: LowerName, - #[derivative(Default(value = "DNS_DEFAULT_ADDRESSES.clone()"))] pub addresses: NameServerAddrGroup, pub listeners: NameServerAddrGroup, } +pub type DnsConfig = ConfigBase; + +impl From for DnsConfig { + fn from(raw: DnsConfigRaw) -> Self { + let default = DnsConfigParsed { + domain: DNS_DEFAULT_DOMAIN.clone(), + name: DNS_DEFAULT_DOMAIN.clone(), + addresses: DNS_DEFAULT_ADDRESSES.clone(), + ..Default::default() + }; + + let parsed = raw.clone().build(default); + Self::new(raw, parsed, ()) + } +} + #[auto_impl::auto_impl(Box, &)] pub trait DnsConfigLoaderExt { fn get_dns(&self) -> DnsConfig; - fn set_dns(&self, dns: Option); + fn set_dns(&self, dns: DnsConfig); } pub type DnsExportConfig = GetExportConfigResponse; diff --git a/easytier/src/dns/config/zone.rs b/easytier/src/dns/config/zone.rs index 9953a399..e2c8fba9 100644 --- a/easytier/src/dns/config/zone.rs +++ b/easytier/src/dns/config/zone.rs @@ -1,47 +1,62 @@ +use crate::common::config::ConfigBase; use crate::dns::config::policy::{DnsExportPolicy, ZonePolicyConfig}; use crate::dns::utils::addr::NameServerAddrGroup; use crate::dns::zone::Zone; use crate::proto::dns::ZoneData; -use derivative::Derivative; -use derive_more::{Deref, Into}; use hickory_proto::rr::LowerName; +use optional_struct::{Applicable, optional_struct}; use serde::{Deserialize, Serialize}; -use std::convert::{TryFrom, TryInto}; +use std::convert::TryFrom; use std::net::{Ipv4Addr, Ipv6Addr}; +use url::Url; -#[derive(Derivative, Debug, Clone, Deserialize, Serialize, Default, Deref, Into)] -#[derivative(PartialEq)] -#[serde(try_from = "ZoneConfigInner", into = "ZoneConfigInner")] -pub struct ZoneConfig { - #[into] - #[derivative(PartialEq = "ignore")] - data: ZoneData, - // User-facing config source of truth used for serde round-trips. - // Keep this in sync with `data` by rebuilding a full ZoneConfig via TryFrom. - // Do not mutate subfields in place and expect `data` to follow. - #[into] - #[deref] - inner: ZoneConfigInner, +#[optional_struct(ZoneConfigRaw)] +#[derive(Debug, Clone, Default, PartialEq, Deserialize, Serialize)] +pub struct ZoneConfigParsed { + #[optional_skip_wrap] + pub origin: LowerName, + pub ttl: u32, + pub records: Vec, + pub forwarders: NameServerAddrGroup, + #[optional_skip_wrap] + #[serde(flatten)] + pub policy: ZonePolicyConfig, + pub fallthrough: bool, } -impl TryFrom for ZoneConfig { +impl From<&ZoneConfigParsed> for ZoneData { + fn from(value: &ZoneConfigParsed) -> Self { + Self::new( + &value.origin, + value.ttl, + &value.records, + value.forwarders.iter().map(Url::from), + value.fallthrough, + ) + } +} + +pub type ZoneConfig = ConfigBase; + +impl TryFrom for ZoneConfig { type Error = anyhow::Error; - fn try_from(value: ZoneConfigInner) -> Result { - // Rebuild both representations together and validate zone semantics. - // Config updates should follow this replacement path. - let data = ZoneData::from(value.clone()); + fn try_from(raw: ZoneConfigRaw) -> Result { + let default = ZoneConfigParsed { + fallthrough: true, + ..Default::default() + }; + + let parsed = raw.clone().build(default); + let data = (&parsed).into(); let _ = Zone::try_from(&data)?; // validation - Ok(Self { data, inner: value }) + + Ok(Self::new(raw, parsed, data)) } } impl ZoneConfig { - pub fn dedicated( - origin: LowerName, - ipv4: Option, - ipv6: Vec, - ) -> anyhow::Result { + pub fn dedicated(origin: LowerName, ipv4: Option, ipv6: Vec) -> Self { let mut records = Vec::new(); if let Some(ipv4) = ipv4 { @@ -55,39 +70,15 @@ impl ZoneConfig { export: Some(DnsExportPolicy::default()), }; - let config = ZoneConfigInner { + let parsed = ZoneConfigParsed { origin, records, policy, ..Default::default() }; - config.try_into() - } -} - -#[derive(Derivative, Debug, Clone, PartialEq, Deserialize, Serialize)] -#[derivative(Default)] -#[serde(default)] -pub struct ZoneConfigInner { - pub origin: LowerName, - pub ttl: u32, - pub records: Vec, - pub forwarders: NameServerAddrGroup, - #[serde(flatten)] - pub policy: ZonePolicyConfig, - #[derivative(Default(value = "true"))] - pub fallthrough: bool, -} - -impl From for ZoneData { - fn from(value: ZoneConfigInner) -> Self { - Self { - origin: value.origin.to_string(), - ttl: value.ttl, - records: value.records, - forwarders: value.forwarders.into(), - fallthrough: value.fallthrough, - } + let data = (&parsed).into(); + + Self::new(Default::default(), parsed, data) } } diff --git a/easytier/src/dns/node_mgr.rs b/easytier/src/dns/node_mgr.rs index 8b2599f6..8f6b2e38 100644 --- a/easytier/src/dns/node_mgr.rs +++ b/easytier/src/dns/node_mgr.rs @@ -451,9 +451,7 @@ mod tests { let zones: Vec<_> = mgr.collect_zones().into_iter().map(Into::into).collect(); let loop_zone = zones .into_iter() - .find(|z: &crate::proto::dns::ZoneData| { - z.origin.trim_end_matches('.') == "filter-loop.test" - }) + .find(|z: &crate::proto::dns::ZoneData| z.content.contains("$ORIGIN filter-loop.test")) .expect("test zone should exist"); let forwarders: HashSet = loop_zone @@ -502,7 +500,7 @@ mod tests { let zone = zones .into_iter() .find(|z: &crate::proto::dns::ZoneData| { - z.origin.trim_end_matches('.') == "cross-node-filter.test" + z.content.contains("$ORIGIN cross-node-filter.test") }) .expect("test zone should exist"); diff --git a/easytier/src/dns/peer_mgr.rs b/easytier/src/dns/peer_mgr.rs index 53870bf1..f676dbf9 100644 --- a/easytier/src/dns/peer_mgr.rs +++ b/easytier/src/dns/peer_mgr.rs @@ -1,5 +1,6 @@ use crate::common::PeerId; use crate::common::global_ctx::ArcGlobalCtx; +use crate::dns::config::zone::ZoneConfig; use crate::dns::config::{ DNS_PEER_REFRESH_ATTEMPTS, DNS_PEER_REFRESH_BACKOFF, DNS_PEER_TTI, DnsExportConfig, DnsGlobalCtxExt, @@ -59,7 +60,7 @@ impl DnsPeerMgrInner { let zones = global_ctx .dns_iter_zones() - .map(Into::into) + .map(ZoneConfig::into_data) .chain( self.peers .iter() @@ -67,7 +68,7 @@ impl DnsPeerMgrInner { ) .collect(); - let config = global_ctx.config.get_dns(); + let config = global_ctx.config.get_dns().into_parsed(); DnsSnapshot { zones, addresses: config.addresses.into(), @@ -256,6 +257,7 @@ mod tests { use std::collections::HashSet; use std::net::Ipv4Addr; use tokio::time::{Duration, sleep}; + use url::Url; async fn create_peer_manager_with_zone( host: &str, @@ -263,17 +265,16 @@ mod tests { record_ip: Ipv4Addr, ) -> Arc { let ctx = get_mock_global_ctx(); - let mut dns = ctx.config.get_dns(); - dns.name = host.parse().unwrap(); - dns.zones.push( - ZoneConfig::dedicated( + let mut dns = ctx.config.get_dns().into_raw(); + dns.name = Some(host.parse().unwrap()); + dns.zones + .get_or_insert_default() + .push(ZoneConfig::dedicated( origin.parse().expect("invalid zone origin"), Some(record_ip), vec![], - ) - .expect("failed to build test zone"), - ); - ctx.config.set_dns(Some(dns)); + )); + ctx.config.set_dns(dns.into()); let (s, _r) = create_packet_recv_chan(); let peer_mgr = Arc::new(PeerManager::new(RouteAlgoType::Ospf, ctx, s)); @@ -295,13 +296,13 @@ mod tests { #[test] fn dns_peer_info_try_from_invalid_zone_rejected() { let cfg = DnsExportConfig { - zones: vec![ZoneData { - origin: "?".to_string(), - ttl: 60, - records: vec!["?".to_string()], - forwarders: vec![], - fallthrough: false, - }], + zones: vec![ZoneData::new( + &".".parse().unwrap(), + 60, + ["?"], + Vec::::new(), + false, + )], }; assert!(DnsPeerInfo::try_from(cfg).is_err()); @@ -333,13 +334,13 @@ mod tests { snapshot .zones .iter() - .any(|z| z.origin.contains("peer-cache.test")) + .any(|z| z.content.contains("$ORIGIN peer-cache.test")) ); assert!( snapshot .zones .iter() - .any(|z| z.origin.contains("local-custom.test")) + .any(|z| z.content.contains("$ORIGIN local-custom.test")) ); } @@ -352,7 +353,7 @@ mod tests { ) .await; let global_ctx = peer_mgr.get_global_ctx(); - let expected = global_ctx.config.get_dns(); + let expected = global_ctx.config.get_dns().into_parsed(); let mgr = DnsPeerMgr::new(peer_mgr, global_ctx); let snapshot = mgr.snapshot(); @@ -417,11 +418,15 @@ mod tests { .await; let snapshot = mgr.snapshot(); - let origins: HashSet<_> = snapshot.zones.into_iter().map(|z| z.origin).collect(); + let contents: HashSet<_> = snapshot.zones.into_iter().map(|z| z.content).collect(); - assert!(origins.iter().any(|z| z.contains("peer-a.test"))); - assert!(origins.iter().any(|z| z.contains("peer-b.test"))); - assert!(origins.iter().any(|z| z.contains("local-multi.test"))); + assert!(contents.iter().any(|z| z.contains("$ORIGIN peer-a.test"))); + assert!(contents.iter().any(|z| z.contains("$ORIGIN peer-b.test"))); + assert!( + contents + .iter() + .any(|z| z.contains("$ORIGIN local-multi.test")) + ); } #[tokio::test] @@ -587,7 +592,7 @@ mod tests { snapshot .zones .iter() - .any(|z| z.origin.contains("remote-export.test")) + .any(|z| z.content.contains("$ORIGIN remote-export.test")) ); } @@ -626,13 +631,13 @@ mod tests { snapshot .zones .iter() - .any(|z| z.origin.contains("remote-a.test")) + .any(|z| z.content.contains("$ORIGIN remote-a.test")) ); assert!( !snapshot .zones .iter() - .any(|z| z.origin.contains("remote-b.test")) + .any(|z| z.content.contains("$ORIGIN remote-b.test")) ); } @@ -794,7 +799,7 @@ mod tests { unchanged_cache .zones .iter() - .any(|z| z.origin.contains("cached-unchanged.test")) + .any(|z| z.content.contains("$ORIGIN cached-unchanged.test")) ); } @@ -825,13 +830,13 @@ mod tests { before .zones .iter() - .any(|z| z.origin.contains("cached-expire.test")) + .any(|z| z.content.contains("$ORIGIN cached-expire.test")) ); assert!( before .zones .iter() - .any(|z| z.origin.contains("local-tti.test")) + .any(|z| z.content.contains("$ORIGIN local-tti.test")) ); let deadline = tokio::time::Instant::now() + DNS_PEER_TTI + Duration::from_secs(3); @@ -840,13 +845,13 @@ mod tests { let expired = !now_snapshot .zones .iter() - .any(|z| z.origin.contains("cached-expire.test")); + .any(|z| z.content.contains("$ORIGIN cached-expire.test")); if expired { assert!( now_snapshot .zones .iter() - .any(|z| z.origin.contains("local-tti.test")) + .any(|z| z.content.contains("local-tti.test")) ); break; } diff --git a/easytier/src/dns/tests.rs b/easytier/src/dns/tests.rs index d99ac521..591f5300 100644 --- a/easytier/src/dns/tests.rs +++ b/easytier/src/dns/tests.rs @@ -9,7 +9,6 @@ use crate::common::config::TomlConfigLoader; use crate::common::global_ctx::GlobalCtx; use crate::common::global_ctx::tests::get_mock_global_ctx; use crate::connector::udp_hole_punch::tests::replace_stun_info_collector; -use crate::dns::config::zone::ZoneConfigInner; use crate::dns::node::DnsNode; use crate::dns::peer_mgr::DnsPeerMgr; use crate::instance::instance::ArcNicCtx; @@ -46,12 +45,12 @@ pub async fn prepare_env_with_tld_dns_zone( ctx.set_hostname(dns_name.to_owned()); ctx.set_ipv4(Some(tun_ip)); - let mut dns_config = ctx.config.get_dns(); - dns_config.name = dns_name.parse().unwrap(); + let mut dns_config = ctx.config.get_dns().into_raw(); + dns_config.name = Some(dns_name.parse().unwrap()); if let Some(zone) = tld_dns_zone { - dns_config.domain = zone.parse().expect("invalid test dns zone"); + dns_config.domain = Some(zone.parse().expect("invalid test dns zone")); } - ctx.config.set_dns(Some(dns_config)); + ctx.config.set_dns(dns_config.into()); let (s, r) = create_packet_recv_chan(); let peer_mgr = Arc::new(PeerManager::new(RouteAlgoType::Ospf, ctx, s)); @@ -105,16 +104,15 @@ pub fn zone_data_a(origin: &str, record: &str) -> ZoneData { } pub fn zone_data_a_with_forwarders(origin: &str, record: &str, forwarders: Vec<&str>) -> ZoneData { - ZoneData { - origin: origin.to_string(), - ttl: 60, - records: vec![format!("@ IN A {record}")], - forwarders: forwarders + ZoneData::new( + &origin.parse().unwrap(), + 60, + [format!("@ IN A {record}")], + forwarders .into_iter() - .map(|f| Url::from_str(f).expect("invalid forwarder")) - .collect(), - fallthrough: false, - } + .map(|f| Url::from_str(f).expect("invalid forwarder")), + false, + ) } pub fn dns_snapshot_with( @@ -349,28 +347,21 @@ async fn wait_peer_zone_visibility( let snapshot = dns.snapshot(); - let visible = snapshot - .zones - .iter() - .any(|z| z.origin.contains(zone_origin_substr)); + let visible = snapshot.zones.iter().any(|z| { + z.content + .contains(&format!("$ORIGIN {}", zone_origin_substr)) + }); if visible == expected_visible { return; } - let origins = snapshot - .zones - .iter() - .map(|z| z.origin.clone()) - .collect::>(); - assert!( Instant::now() < deadline, - "zone visibility mismatch for '{}': expected {}, got {}, current origins: {:?}", + "zone visibility mismatch for '{}': expected {}, got {}", zone_origin_substr, expected_visible, visible, - origins ); tokio::time::sleep(Duration::from_millis(200)).await; } @@ -504,7 +495,7 @@ records = ["secret IN A 10.99.0.9"] !snapshot .zones .iter() - .any(|z| z.origin.contains("private.mesh-test")), + .any(|z| z.content.contains("$ORIGIN private.mesh-test")), "zone without [dns.zone.export] should not be exported to peer snapshot" ); } @@ -552,7 +543,7 @@ disabled = true !snapshot .zones .iter() - .any(|z| z.origin.contains("disabled.mesh-test")), + .any(|z| z.content.contains("$ORIGIN disabled.mesh-test")), "zone with [dns.zone.export] disabled=true should not be exported" ); } @@ -580,16 +571,17 @@ records = ["api IN A 10.80.0.1"] check_dns_record_at(server_addr, "api.patch.mesh-test.", "10.80.0.1").await; - let mut dns = peer.get_global_ctx().config.get_dns(); - let zone_idx = dns - .zones + let mut dns = peer.get_global_ctx().config.get_dns().into_raw(); + let mut zones = dns.zones.unwrap(); + let zone_idx = zones .iter() .position(|z| z.origin.to_string().contains("patch.mesh-test")) .expect("patch zone should exist"); - let mut zone: ZoneConfigInner = dns.zones[zone_idx].clone().into(); - zone.records = vec!["api IN A 10.80.0.2".to_string()]; - dns.zones[zone_idx] = zone.try_into().expect("patch zone update should be valid"); - peer.get_global_ctx().config.set_dns(Some(dns)); + let mut zone = zones[zone_idx].clone().into_raw(); + zone.records = Some(vec!["api IN A 10.80.0.2".to_string()]); + zones[zone_idx] = zone.try_into().expect("patch zone update should be valid"); + dns.zones = Some(zones); + peer.get_global_ctx().config.set_dns(dns.into()); peer.get_global_ctx() .issue_event(crate::common::global_ctx::GlobalCtxEvent::ConfigPatched( crate::proto::api::config::InstanceConfigPatch::default(), @@ -619,14 +611,16 @@ async fn config_patch_reloads_listener_binding() { let new_addr = SocketAddr::new(Ipv4Addr::LOCALHOST.into(), listener_new); check_dns_record_at(old_addr, "listener-patch.mesh-test.", "10.144.150.11").await; - let mut dns = peer.get_global_ctx().config.get_dns(); - dns.listeners = vec![ - format!("udp://127.0.0.1:{listener_new}") - .parse() - .expect("invalid listener"), - ] - .into(); - peer.get_global_ctx().config.set_dns(Some(dns)); + let mut dns = peer.get_global_ctx().config.get_dns().into_raw(); + dns.listeners = Some( + vec![ + format!("udp://127.0.0.1:{listener_new}") + .parse() + .expect("invalid listener"), + ] + .into(), + ); + peer.get_global_ctx().config.set_dns(dns.into()); peer.get_global_ctx() .issue_event(crate::common::global_ctx::GlobalCtxEvent::ConfigPatched( crate::proto::api::config::InstanceConfigPatch::default(), @@ -894,7 +888,7 @@ records = ["secret IN A 10.66.3.8"] !snapshot .zones .iter() - .any(|z| z.origin.contains("private-c.mesh5-test")), + .any(|z| z.content.contains("$ORIGIN private-c.mesh5-test")), "zone without [dns.zone.export] should not sync over multi-hop" ); } @@ -929,7 +923,7 @@ async fn config_string_two_nodes_peer_dns_offline_then_rejoin() { .snapshot() .zones .iter() - .any(|z| z.origin.contains("node-b6.mesh6-test")), + .any(|z| z.content.contains("$ORIGIN node-b6.mesh6-test")), "peer B self zone should be visible after initial refresh" ); @@ -966,7 +960,7 @@ async fn config_string_two_nodes_peer_dns_offline_then_rejoin() { .snapshot() .zones .iter() - .any(|z| z.origin.contains("node-b6.mesh6-test")), + .any(|z| z.content.contains("$ORIGIN node-b6.mesh6-test")), "peer B self zone should disappear after route withdrawal and cache expiry" ); @@ -985,7 +979,7 @@ async fn config_string_two_nodes_peer_dns_offline_then_rejoin() { .snapshot() .zones .iter() - .any(|z| z.origin.contains("node-b6.mesh6-test")), + .any(|z| z.content.contains("$ORIGIN node-b6.mesh6-test")), "peer B self zone should be restored after DNS RPC rejoins" ); } diff --git a/easytier/src/dns/utils/addr.rs b/easytier/src/dns/utils/addr.rs index c8d3fdff..9b2a6783 100644 --- a/easytier/src/dns/utils/addr.rs +++ b/easytier/src/dns/utils/addr.rs @@ -39,9 +39,15 @@ impl From<(IpAddr, &ConnectionConfig)> for NameServerAddr { } } +impl From<&NameServerAddr> for Url { + fn from(value: &NameServerAddr) -> Self { + Url::parse(&format!("{}://{}", value.protocol, value.addr)).unwrap() + } +} + impl From for Url { fn from(value: NameServerAddr) -> Self { - Url::parse(&format!("{}://{}", value.protocol, value.addr)).unwrap() + (&value).into() } } diff --git a/easytier/src/dns/zone.rs b/easytier/src/dns/zone.rs index c40ca507..bda697c0 100644 --- a/easytier/src/dns/zone.rs +++ b/easytier/src/dns/zone.rs @@ -1,6 +1,6 @@ use crate::dns::utils::addr::{NameServerAddr, NameServerAddrGroup}; use crate::dns::utils::zone_handler::{ArcZoneHandler, ChainedZoneHandler}; -use crate::proto; +use crate::proto::dns::ZoneData; use crate::proto::utils::RepeatedMessageModel; use crate::utils::dns::resolver_conf; use hickory_net::runtime::TokioRuntimeProvider; @@ -13,6 +13,7 @@ use indexmap::IndexMap; use itertools::chain; use std::collections::BTreeMap; use std::sync::Arc; +use url::Url; #[derive(Debug, Clone)] pub struct Zone { @@ -89,11 +90,11 @@ impl Zone { } } -impl TryFrom<&proto::dns::ZoneData> for Zone { +impl TryFrom<&ZoneData> for Zone { type Error = anyhow::Error; - fn try_from(value: &proto::dns::ZoneData) -> Result { - let (origin, records) = Parser::new(value.to_string(), None, None) + fn try_from(value: &ZoneData) -> Result { + let (origin, records) = Parser::new(&value.content, None, None) .parse() .map_err(|e| anyhow::anyhow!("failed to parse zone data: {e}"))?; @@ -117,14 +118,13 @@ impl TryFrom<&proto::dns::ZoneData> for Zone { } } -impl From for proto::dns::ZoneData { +impl From for ZoneData { fn from(value: Zone) -> Self { let records = value .records .values() .flat_map(RecordSet::records_without_rrsigs) - .map(ToString::to_string) - .collect(); + .map(ToString::to_string); let forwarders = value .forward @@ -132,16 +132,9 @@ impl From for proto::dns::ZoneData { .flat_map(|f| f.name_servers.into_iter()) .map(|ns| (&ns).into()) .flat_map(NameServerAddrGroup::into_iter) - .map(Into::into) - .collect(); + .map(Url::from); - Self { - origin: value.origin.to_string(), - ttl: 0, - records, - forwarders, - fallthrough: value.fallthrough, - } + Self::new(&value.origin, 0, records, forwarders, value.fallthrough) } } @@ -202,16 +195,15 @@ mod tests { forwarders: Vec<&str>, fallthrough: bool, ) -> ZoneData { - ZoneData { - origin: origin.to_string(), - ttl: 60, - records: records.into_iter().map(ToString::to_string).collect(), - forwarders: forwarders - .into_iter() - .map(|f| Url::from_str(f).expect("invalid forwarder")) - .collect(), + ZoneData::new( + &origin.parse().unwrap(), + 60, + records, + forwarders.into_iter().map(|url| Url { + url: url.to_string(), + }), fallthrough, - } + ) } fn zone_data(origin: &str, records: Vec<&str>, forwarders: Vec<&str>) -> ZoneData { @@ -313,10 +305,6 @@ mod tests { assert_eq!(zone.forward.as_ref().unwrap().name_servers.len(), 2); let serialized = ZoneData::from(zone.clone()); - assert_eq!(serialized.origin, "roundtrip.test."); - assert_eq!(serialized.records.len(), 2); - assert_eq!(serialized.forwarders.len(), 2); - let reparsed = Zone::try_from(&serialized)?; assert_eq!(reparsed.origin.to_string(), "roundtrip.test."); assert_eq!(reparsed.iter_records().count(), 2); diff --git a/easytier/src/proto/dns.proto b/easytier/src/proto/dns.proto index 68b0fbd2..1bf0dd7f 100644 --- a/easytier/src/proto/dns.proto +++ b/easytier/src/proto/dns.proto @@ -5,11 +5,9 @@ import "common.proto"; package dns; message ZoneData { - string origin = 1; - uint32 ttl = 2; - repeated string records = 3; - repeated common.Url forwarders = 4; - bool fallthrough = 5; + string content = 1; + repeated common.Url forwarders = 2; + bool fallthrough = 3; } message GetExportConfigRequest {} diff --git a/easytier/src/proto/dns.rs b/easytier/src/proto/dns.rs index c05969dc..cbf6e17a 100644 --- a/easytier/src/proto/dns.rs +++ b/easytier/src/proto/dns.rs @@ -1,5 +1,7 @@ +use crate::proto::common::Url; use crate::proto::utils::TransientDigest; -use std::fmt::Display; +use hickory_proto::rr::LowerName; +use std::fmt::Write; include!(concat!(env!("OUT_DIR"), "/dns.rs")); @@ -10,31 +12,44 @@ impl HeartbeatRequest { } } -impl Display for ZoneData { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - writeln!(f, "; EasyTier Magic DNS zone data")?; - writeln!(f, "; https://github.com/easytier/easytier")?; +impl ZoneData { + pub fn new( + origin: &LowerName, + ttl: u32, + records: Records, + forwarders: Urls, + fallthrough: bool, + ) -> Self + where + Records: IntoIterator, + R: AsRef, + Urls: IntoIterator, + U: Into, + { + let mut content = String::new(); - if !self.forwarders.is_empty() { - writeln!(f, "; Forwarders:")?; - for forwarder in &self.forwarders { - writeln!(f, "; \t{}", forwarder)?; - } - } - writeln!(f)?; + content.push_str("; EasyTier Magic DNS zone data\n"); + content.push_str("; https://github.com/easytier/easytier\n"); - write!(f, "$ORIGIN {}", self.origin)?; - if !self.origin.ends_with('.') { - write!(f, ".")?; - } - writeln!(f)?; - - writeln!(f, "$TTL {}", self.ttl)?; - - for record in &self.records { - writeln!(f, "{}", record)?; + let mut origin = origin.to_string(); + if !origin.ends_with('.') { + origin.push('.'); } - Ok(()) + writeln!(content, "$ORIGIN {}", origin).unwrap(); + writeln!(content, "$TTL {}", ttl).unwrap(); + + for record in records { + content.push_str(record.as_ref()); + content.push('\n'); + } + + let forwarders = forwarders.into_iter().map(Into::into).collect(); + + Self { + content, + forwarders, + fallthrough, + } } } diff --git a/easytier/src/tests/three_node.rs b/easytier/src/tests/three_node.rs index 58fb6b3e..fd08cb60 100644 --- a/easytier/src/tests/three_node.rs +++ b/easytier/src/tests/three_node.rs @@ -4127,14 +4127,14 @@ pub async fn three_node_dns_export() { "tcp", |cfg| { use crate::dns::config::zone::ZoneConfig; - use crate::dns::config::{DnsConfig, DnsConfigLoaderExt}; + use crate::dns::config::{DnsConfigLoaderExt, DnsConfigRaw}; use hickory_proto::rr::LowerName; use std::str::FromStr; let inst_name = cfg.get_inst_name(); let origin = LowerName::from_str(&format!("{}.com.", inst_name)).unwrap(); - let mut dns_config = DnsConfig { - name: LowerName::from_str(&inst_name).unwrap(), + let mut dns_config = DnsConfigRaw { + name: Some(LowerName::from_str(&inst_name).unwrap()), ..Default::default() }; @@ -4147,7 +4147,8 @@ pub async fn three_node_dns_export() { dns_config .zones - .push(ZoneConfig::dedicated(origin, ipv4, vec![]).unwrap()); + .get_or_insert_default() + .push(ZoneConfig::dedicated(origin, ipv4, vec![])); let listener_port = match inst_name.as_str() { "inst1" => 5351, @@ -4155,14 +4156,16 @@ pub async fn three_node_dns_export() { "inst3" => 5353, _ => 5350, }; - dns_config.listeners = vec![ - format!("udp://127.0.0.1:{}", listener_port) - .parse() - .unwrap(), - ] - .into(); + dns_config.listeners = Some( + vec![ + format!("udp://127.0.0.1:{}", listener_port) + .parse() + .unwrap(), + ] + .into(), + ); - cfg.set_dns(Some(dns_config)); + cfg.set_dns(dns_config.into()); cfg }, false, @@ -4195,14 +4198,14 @@ pub async fn three_node_dns_export_chain() { let cfg_cb = |cfg: TomlConfigLoader| { use crate::dns::config::zone::ZoneConfig; - use crate::dns::config::{DnsConfig, DnsConfigLoaderExt}; + use crate::dns::config::{DnsConfigLoaderExt, DnsConfigRaw}; use hickory_proto::rr::LowerName; use std::str::FromStr; let inst_name = cfg.get_inst_name(); let origin = LowerName::from_str(&format!("{}.com.", inst_name)).unwrap(); - let mut dns_config = DnsConfig { - name: LowerName::from_str(&inst_name).unwrap(), + let mut dns_config = DnsConfigRaw { + name: Some(LowerName::from_str(&inst_name).unwrap()), ..Default::default() }; @@ -4215,7 +4218,8 @@ pub async fn three_node_dns_export_chain() { dns_config .zones - .push(ZoneConfig::dedicated(origin, ipv4, vec![]).unwrap()); + .get_or_insert_default() + .push(ZoneConfig::dedicated(origin, ipv4, vec![])); let listener_port = match inst_name.as_str() { "inst1" => 5351, @@ -4223,14 +4227,16 @@ pub async fn three_node_dns_export_chain() { "inst3" => 5353, _ => 5350, }; - dns_config.listeners = vec![ - format!("udp://127.0.0.1:{}", listener_port) - .parse() - .unwrap(), - ] - .into(); + dns_config.listeners = Some( + vec![ + format!("udp://127.0.0.1:{}", listener_port) + .parse() + .unwrap(), + ] + .into(), + ); - cfg.set_dns(Some(dns_config)); + cfg.set_dns(dns_config.into()); let mut flags = cfg.get_flags(); flags.disable_p2p = true;