config: rewrite

test

test

origin

test

test
This commit is contained in:
Luna Yao
2026-04-29 03:33:32 +02:00
parent 0a23546e3f
commit 92af568d8b
13 changed files with 338 additions and 251 deletions
+2
View File
@@ -50,6 +50,8 @@ time = "0.3"
toml = "0.8.12" toml = "0.8.12"
chrono = { version = "0.4.37", features = ["serde"] } chrono = { version = "0.4.37", features = ["serde"] }
optional_struct = "0.5.2"
guarden = "0.1" guarden = "0.1"
delegate = "0.13.5" delegate = "0.13.5"
+87 -18
View File
@@ -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 super::env_parser;
use crate::utils::dns::sanitize; use crate::utils::dns::sanitize;
use crate::{ use crate::{
@@ -25,6 +10,89 @@ use crate::{
tunnel::{IpScheme, TunnelScheme, generate_digest_from_str}, tunnel::{IpScheme, TunnelScheme, generate_digest_from_str},
utils, 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>, <ConfigBase<Raw, Parsed, Data> as TryFrom<Raw>>::Error: Display"
)]
pub struct ConfigBase<Raw, Parsed, Data>
where
Raw: Applicable<Base = Parsed>,
ConfigBase<Raw, Parsed, Data>: TryFrom<Raw, Error: Debug>,
{
#[derivative(PartialEq = "ignore")]
raw: Raw,
#[deref]
parsed: Parsed,
#[derivative(PartialEq = "ignore")]
data: Data,
}
impl<Raw, Parsed, Data> Serialize for ConfigBase<Raw, Parsed, Data>
where
Raw: Applicable<Base = Parsed> + Serialize,
ConfigBase<Raw, Parsed, Data>: TryFrom<Raw, Error: Debug>,
{
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
self.raw.serialize(serializer)
}
}
impl<Raw, Parsed, Data> Default for ConfigBase<Raw, Parsed, Data>
where
Raw: Applicable<Base = Parsed> + Default,
ConfigBase<Raw, Parsed, Data>: TryFrom<Raw, Error: Debug>,
{
fn default() -> Self {
Raw::default().try_into().unwrap()
}
}
impl<Raw, Parsed, Data> ConfigBase<Raw, Parsed, Data>
where
Raw: Applicable<Base = Parsed>,
ConfigBase<Raw, Parsed, Data>: TryFrom<Raw, Error: Debug>,
{
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<Self, <Self as TryFrom<Raw>>::Error> {
let mut raw = self.into_raw();
config.apply_to_opt(&mut raw);
raw.try_into()
}
}
pub type Flags = crate::proto::common::FlagsInConfig; pub type Flags = crate::proto::common::FlagsInConfig;
@@ -554,7 +622,8 @@ struct Config {
proxy_network: Option<Vec<ProxyNetworkConfig>>, proxy_network: Option<Vec<ProxyNetworkConfig>>,
#[cfg(feature = "magic-dns")] #[cfg(feature = "magic-dns")]
dns: Option<DnsConfig>, #[serde(default)]
dns: DnsConfig,
vpn_portal_config: Option<VpnPortalConfig>, vpn_portal_config: Option<VpnPortalConfig>,
@@ -655,9 +724,9 @@ impl DnsConfigLoaderExt for TomlConfigLoader {
cfg_select! { cfg_select! {
feature = "magic-dns" => { feature = "magic-dns" => {
fn get_dns(&self) -> DnsConfig { 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<DnsConfig>) { fn set_dns(&self, config: DnsConfig) {
self.config.lock().unwrap().dns = config; self.config.lock().unwrap().dns = config;
} }
} }
+7 -7
View File
@@ -21,7 +21,7 @@ use super::{
stun::{StunInfoCollector, StunInfoCollectorTrait}, stun::{StunInfoCollector, StunInfoCollectorTrait},
}; };
#[cfg(feature = "magic-dns")] #[cfg(feature = "magic-dns")]
use crate::dns::config::{DnsConfigLoaderExt, DnsExportConfig, DnsGlobalCtxExt}; use crate::dns::config::{DnsConfigLoaderExt, DnsExportConfig, DnsGlobalCtxExt, zone::ZoneConfig};
use crate::{ use crate::{
common::{ common::{
config::ProxyNetworkConfig, shrink_dashmap, stats_manager::StatsManager, config::ProxyNetworkConfig, shrink_dashmap, stats_manager::StatsManager,
@@ -461,7 +461,7 @@ impl GlobalCtx {
} }
pub fn get_hostname(&self) -> String { pub fn get_hostname(&self) -> String {
return self.hostname.lock().unwrap().clone(); self.hostname.lock().unwrap().clone()
} }
pub fn set_hostname(&self, hostname: String) { pub fn set_hostname(&self, hostname: String) {
@@ -713,7 +713,7 @@ impl GlobalCtx {
#[cfg(feature = "magic-dns")] #[cfg(feature = "magic-dns")]
impl DnsGlobalCtxExt for GlobalCtx { 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 dns = self.config.get_dns();
let mut hostname = dns.name.to_string(); let mut hostname = dns.name.to_string();
if hostname.is_empty() { if hostname.is_empty() {
@@ -728,7 +728,7 @@ impl DnsGlobalCtxExt for GlobalCtx {
let ipv6 = self.get_ipv6().map(|ip| ip.address()); let ipv6 = self.get_ipv6().map(|ip| ip.address());
let ipv6 = ipv6.map(|a| vec![a]).unwrap_or_default(); 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 { fn dns_export_config(&self) -> DnsExportConfig {
@@ -736,13 +736,13 @@ impl DnsGlobalCtxExt for GlobalCtx {
zones: self zones: self
.dns_iter_zones() .dns_iter_zones()
.filter(|z| z.policy.export.as_ref().is_some_and(|f| !f.disabled)) // TODO: check policies of parent 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(), .collect(),
} }
} }
fn dns_iter_zones(&self) -> impl Iterator<Item = crate::dns::config::zone::ZoneConfig> { fn dns_iter_zones(&self) -> impl Iterator<Item = ZoneConfig> {
iter::once(self.dns_self_zone()).chain(self.config.get_dns().zones) iter::once(self.dns_self_zone()).chain(self.config.get_dns().into_parsed().zones)
} }
} }
+23 -8
View File
@@ -1,33 +1,48 @@
use crate::common::config::ConfigBase;
use crate::dns::config::policy::DnsPolicyConfig; use crate::dns::config::policy::DnsPolicyConfig;
use crate::dns::config::zone::ZoneConfig; use crate::dns::config::zone::ZoneConfig;
use crate::dns::config::{DNS_DEFAULT_ADDRESSES, DNS_DEFAULT_DOMAIN}; use crate::dns::config::{DNS_DEFAULT_ADDRESSES, DNS_DEFAULT_DOMAIN};
use crate::dns::utils::addr::NameServerAddrGroup; use crate::dns::utils::addr::NameServerAddrGroup;
use crate::proto::dns::GetExportConfigResponse; use crate::proto::dns::GetExportConfigResponse;
use derivative::Derivative;
use hickory_proto::rr::LowerName; use hickory_proto::rr::LowerName;
use optional_struct::{Applicable, optional_struct};
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use std::collections::HashMap; use std::collections::HashMap;
#[derive(Derivative, Debug, Clone, Deserialize, Serialize, PartialEq)] #[optional_struct(DnsConfigRaw)]
#[derivative(Default)] #[derive(Debug, Clone, Default, PartialEq, Deserialize, Serialize)]
#[serde(default)] pub struct DnsConfigParsed {
pub struct DnsConfig {
#[serde(rename = "zone")] #[serde(rename = "zone")]
pub zones: Vec<ZoneConfig>, pub zones: Vec<ZoneConfig>,
#[optional_skip_wrap]
#[serde(flatten)] #[serde(flatten)]
pub policies: HashMap<LowerName, DnsPolicyConfig>, pub policies: HashMap<LowerName, DnsPolicyConfig>,
pub name: LowerName, pub name: LowerName,
#[derivative(Default(value = "DNS_DEFAULT_DOMAIN.clone()"))]
pub domain: LowerName, pub domain: LowerName,
#[derivative(Default(value = "DNS_DEFAULT_ADDRESSES.clone()"))]
pub addresses: NameServerAddrGroup, pub addresses: NameServerAddrGroup,
pub listeners: NameServerAddrGroup, pub listeners: NameServerAddrGroup,
} }
pub type DnsConfig = ConfigBase<DnsConfigRaw, DnsConfigParsed, ()>;
impl From<DnsConfigRaw> 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, &)] #[auto_impl::auto_impl(Box, &)]
pub trait DnsConfigLoaderExt { pub trait DnsConfigLoaderExt {
fn get_dns(&self) -> DnsConfig; fn get_dns(&self) -> DnsConfig;
fn set_dns(&self, dns: Option<DnsConfig>); fn set_dns(&self, dns: DnsConfig);
} }
pub type DnsExportConfig = GetExportConfigResponse; pub type DnsExportConfig = GetExportConfigResponse;
+46 -55
View File
@@ -1,47 +1,62 @@
use crate::common::config::ConfigBase;
use crate::dns::config::policy::{DnsExportPolicy, ZonePolicyConfig}; use crate::dns::config::policy::{DnsExportPolicy, ZonePolicyConfig};
use crate::dns::utils::addr::NameServerAddrGroup; use crate::dns::utils::addr::NameServerAddrGroup;
use crate::dns::zone::Zone; use crate::dns::zone::Zone;
use crate::proto::dns::ZoneData; use crate::proto::dns::ZoneData;
use derivative::Derivative;
use derive_more::{Deref, Into};
use hickory_proto::rr::LowerName; use hickory_proto::rr::LowerName;
use optional_struct::{Applicable, optional_struct};
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use std::convert::{TryFrom, TryInto}; use std::convert::TryFrom;
use std::net::{Ipv4Addr, Ipv6Addr}; use std::net::{Ipv4Addr, Ipv6Addr};
use url::Url;
#[derive(Derivative, Debug, Clone, Deserialize, Serialize, Default, Deref, Into)] #[optional_struct(ZoneConfigRaw)]
#[derivative(PartialEq)] #[derive(Debug, Clone, Default, PartialEq, Deserialize, Serialize)]
#[serde(try_from = "ZoneConfigInner", into = "ZoneConfigInner")] pub struct ZoneConfigParsed {
pub struct ZoneConfig { #[optional_skip_wrap]
#[into] pub origin: LowerName,
#[derivative(PartialEq = "ignore")] pub ttl: u32,
data: ZoneData, pub records: Vec<String>,
// User-facing config source of truth used for serde round-trips. pub forwarders: NameServerAddrGroup,
// Keep this in sync with `data` by rebuilding a full ZoneConfig via TryFrom. #[optional_skip_wrap]
// Do not mutate subfields in place and expect `data` to follow. #[serde(flatten)]
#[into] pub policy: ZonePolicyConfig,
#[deref] pub fallthrough: bool,
inner: ZoneConfigInner,
} }
impl TryFrom<ZoneConfigInner> 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<ZoneConfigRaw, ZoneConfigParsed, ZoneData>;
impl TryFrom<ZoneConfigRaw> for ZoneConfig {
type Error = anyhow::Error; type Error = anyhow::Error;
fn try_from(value: ZoneConfigInner) -> Result<Self, Self::Error> { fn try_from(raw: ZoneConfigRaw) -> Result<Self, Self::Error> {
// Rebuild both representations together and validate zone semantics. let default = ZoneConfigParsed {
// Config updates should follow this replacement path. fallthrough: true,
let data = ZoneData::from(value.clone()); ..Default::default()
};
let parsed = raw.clone().build(default);
let data = (&parsed).into();
let _ = Zone::try_from(&data)?; // validation let _ = Zone::try_from(&data)?; // validation
Ok(Self { data, inner: value })
Ok(Self::new(raw, parsed, data))
} }
} }
impl ZoneConfig { impl ZoneConfig {
pub fn dedicated( pub fn dedicated(origin: LowerName, ipv4: Option<Ipv4Addr>, ipv6: Vec<Ipv6Addr>) -> Self {
origin: LowerName,
ipv4: Option<Ipv4Addr>,
ipv6: Vec<Ipv6Addr>,
) -> anyhow::Result<Self> {
let mut records = Vec::new(); let mut records = Vec::new();
if let Some(ipv4) = ipv4 { if let Some(ipv4) = ipv4 {
@@ -55,39 +70,15 @@ impl ZoneConfig {
export: Some(DnsExportPolicy::default()), export: Some(DnsExportPolicy::default()),
}; };
let config = ZoneConfigInner { let parsed = ZoneConfigParsed {
origin, origin,
records, records,
policy, policy,
..Default::default() ..Default::default()
}; };
config.try_into() let data = (&parsed).into();
}
} Self::new(Default::default(), parsed, data)
#[derive(Derivative, Debug, Clone, PartialEq, Deserialize, Serialize)]
#[derivative(Default)]
#[serde(default)]
pub struct ZoneConfigInner {
pub origin: LowerName,
pub ttl: u32,
pub records: Vec<String>,
pub forwarders: NameServerAddrGroup,
#[serde(flatten)]
pub policy: ZonePolicyConfig,
#[derivative(Default(value = "true"))]
pub fallthrough: bool,
}
impl From<ZoneConfigInner> 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,
}
} }
} }
+2 -4
View File
@@ -451,9 +451,7 @@ mod tests {
let zones: Vec<_> = mgr.collect_zones().into_iter().map(Into::into).collect(); let zones: Vec<_> = mgr.collect_zones().into_iter().map(Into::into).collect();
let loop_zone = zones let loop_zone = zones
.into_iter() .into_iter()
.find(|z: &crate::proto::dns::ZoneData| { .find(|z: &crate::proto::dns::ZoneData| z.content.contains("$ORIGIN filter-loop.test"))
z.origin.trim_end_matches('.') == "filter-loop.test"
})
.expect("test zone should exist"); .expect("test zone should exist");
let forwarders: HashSet<NameServerAddr> = loop_zone let forwarders: HashSet<NameServerAddr> = loop_zone
@@ -502,7 +500,7 @@ mod tests {
let zone = zones let zone = zones
.into_iter() .into_iter()
.find(|z: &crate::proto::dns::ZoneData| { .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"); .expect("test zone should exist");
+37 -32
View File
@@ -1,5 +1,6 @@
use crate::common::PeerId; use crate::common::PeerId;
use crate::common::global_ctx::ArcGlobalCtx; use crate::common::global_ctx::ArcGlobalCtx;
use crate::dns::config::zone::ZoneConfig;
use crate::dns::config::{ use crate::dns::config::{
DNS_PEER_REFRESH_ATTEMPTS, DNS_PEER_REFRESH_BACKOFF, DNS_PEER_TTI, DnsExportConfig, DNS_PEER_REFRESH_ATTEMPTS, DNS_PEER_REFRESH_BACKOFF, DNS_PEER_TTI, DnsExportConfig,
DnsGlobalCtxExt, DnsGlobalCtxExt,
@@ -59,7 +60,7 @@ impl DnsPeerMgrInner {
let zones = global_ctx let zones = global_ctx
.dns_iter_zones() .dns_iter_zones()
.map(Into::into) .map(ZoneConfig::into_data)
.chain( .chain(
self.peers self.peers
.iter() .iter()
@@ -67,7 +68,7 @@ impl DnsPeerMgrInner {
) )
.collect(); .collect();
let config = global_ctx.config.get_dns(); let config = global_ctx.config.get_dns().into_parsed();
DnsSnapshot { DnsSnapshot {
zones, zones,
addresses: config.addresses.into(), addresses: config.addresses.into(),
@@ -256,6 +257,7 @@ mod tests {
use std::collections::HashSet; use std::collections::HashSet;
use std::net::Ipv4Addr; use std::net::Ipv4Addr;
use tokio::time::{Duration, sleep}; use tokio::time::{Duration, sleep};
use url::Url;
async fn create_peer_manager_with_zone( async fn create_peer_manager_with_zone(
host: &str, host: &str,
@@ -263,17 +265,16 @@ mod tests {
record_ip: Ipv4Addr, record_ip: Ipv4Addr,
) -> Arc<PeerManager> { ) -> Arc<PeerManager> {
let ctx = get_mock_global_ctx(); let ctx = get_mock_global_ctx();
let mut dns = ctx.config.get_dns(); let mut dns = ctx.config.get_dns().into_raw();
dns.name = host.parse().unwrap(); dns.name = Some(host.parse().unwrap());
dns.zones.push( dns.zones
ZoneConfig::dedicated( .get_or_insert_default()
.push(ZoneConfig::dedicated(
origin.parse().expect("invalid zone origin"), origin.parse().expect("invalid zone origin"),
Some(record_ip), Some(record_ip),
vec![], vec![],
) ));
.expect("failed to build test zone"), ctx.config.set_dns(dns.into());
);
ctx.config.set_dns(Some(dns));
let (s, _r) = create_packet_recv_chan(); let (s, _r) = create_packet_recv_chan();
let peer_mgr = Arc::new(PeerManager::new(RouteAlgoType::Ospf, ctx, s)); let peer_mgr = Arc::new(PeerManager::new(RouteAlgoType::Ospf, ctx, s));
@@ -295,13 +296,13 @@ mod tests {
#[test] #[test]
fn dns_peer_info_try_from_invalid_zone_rejected() { fn dns_peer_info_try_from_invalid_zone_rejected() {
let cfg = DnsExportConfig { let cfg = DnsExportConfig {
zones: vec![ZoneData { zones: vec![ZoneData::new(
origin: "?".to_string(), &".".parse().unwrap(),
ttl: 60, 60,
records: vec!["?".to_string()], ["?"],
forwarders: vec![], Vec::<Url>::new(),
fallthrough: false, false,
}], )],
}; };
assert!(DnsPeerInfo::try_from(cfg).is_err()); assert!(DnsPeerInfo::try_from(cfg).is_err());
@@ -333,13 +334,13 @@ mod tests {
snapshot snapshot
.zones .zones
.iter() .iter()
.any(|z| z.origin.contains("peer-cache.test")) .any(|z| z.content.contains("$ORIGIN peer-cache.test"))
); );
assert!( assert!(
snapshot snapshot
.zones .zones
.iter() .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; .await;
let global_ctx = peer_mgr.get_global_ctx(); 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 mgr = DnsPeerMgr::new(peer_mgr, global_ctx);
let snapshot = mgr.snapshot(); let snapshot = mgr.snapshot();
@@ -417,11 +418,15 @@ mod tests {
.await; .await;
let snapshot = mgr.snapshot(); 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!(contents.iter().any(|z| z.contains("$ORIGIN peer-a.test")));
assert!(origins.iter().any(|z| z.contains("peer-b.test"))); assert!(contents.iter().any(|z| z.contains("$ORIGIN peer-b.test")));
assert!(origins.iter().any(|z| z.contains("local-multi.test"))); assert!(
contents
.iter()
.any(|z| z.contains("$ORIGIN local-multi.test"))
);
} }
#[tokio::test] #[tokio::test]
@@ -587,7 +592,7 @@ mod tests {
snapshot snapshot
.zones .zones
.iter() .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 snapshot
.zones .zones
.iter() .iter()
.any(|z| z.origin.contains("remote-a.test")) .any(|z| z.content.contains("$ORIGIN remote-a.test"))
); );
assert!( assert!(
!snapshot !snapshot
.zones .zones
.iter() .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 unchanged_cache
.zones .zones
.iter() .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 before
.zones .zones
.iter() .iter()
.any(|z| z.origin.contains("cached-expire.test")) .any(|z| z.content.contains("$ORIGIN cached-expire.test"))
); );
assert!( assert!(
before before
.zones .zones
.iter() .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); let deadline = tokio::time::Instant::now() + DNS_PEER_TTI + Duration::from_secs(3);
@@ -840,13 +845,13 @@ mod tests {
let expired = !now_snapshot let expired = !now_snapshot
.zones .zones
.iter() .iter()
.any(|z| z.origin.contains("cached-expire.test")); .any(|z| z.content.contains("$ORIGIN cached-expire.test"));
if expired { if expired {
assert!( assert!(
now_snapshot now_snapshot
.zones .zones
.iter() .iter()
.any(|z| z.origin.contains("local-tti.test")) .any(|z| z.content.contains("local-tti.test"))
); );
break; break;
} }
+41 -47
View File
@@ -9,7 +9,6 @@ use crate::common::config::TomlConfigLoader;
use crate::common::global_ctx::GlobalCtx; use crate::common::global_ctx::GlobalCtx;
use crate::common::global_ctx::tests::get_mock_global_ctx; use crate::common::global_ctx::tests::get_mock_global_ctx;
use crate::connector::udp_hole_punch::tests::replace_stun_info_collector; 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::node::DnsNode;
use crate::dns::peer_mgr::DnsPeerMgr; use crate::dns::peer_mgr::DnsPeerMgr;
use crate::instance::instance::ArcNicCtx; 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_hostname(dns_name.to_owned());
ctx.set_ipv4(Some(tun_ip)); ctx.set_ipv4(Some(tun_ip));
let mut dns_config = ctx.config.get_dns(); let mut dns_config = ctx.config.get_dns().into_raw();
dns_config.name = dns_name.parse().unwrap(); dns_config.name = Some(dns_name.parse().unwrap());
if let Some(zone) = tld_dns_zone { 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 (s, r) = create_packet_recv_chan();
let peer_mgr = Arc::new(PeerManager::new(RouteAlgoType::Ospf, ctx, s)); 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 { pub fn zone_data_a_with_forwarders(origin: &str, record: &str, forwarders: Vec<&str>) -> ZoneData {
ZoneData { ZoneData::new(
origin: origin.to_string(), &origin.parse().unwrap(),
ttl: 60, 60,
records: vec![format!("@ IN A {record}")], [format!("@ IN A {record}")],
forwarders: forwarders forwarders
.into_iter() .into_iter()
.map(|f| Url::from_str(f).expect("invalid forwarder")) .map(|f| Url::from_str(f).expect("invalid forwarder")),
.collect(), false,
fallthrough: false, )
}
} }
pub fn dns_snapshot_with( pub fn dns_snapshot_with(
@@ -349,28 +347,21 @@ async fn wait_peer_zone_visibility(
let snapshot = dns.snapshot(); let snapshot = dns.snapshot();
let visible = snapshot let visible = snapshot.zones.iter().any(|z| {
.zones z.content
.iter() .contains(&format!("$ORIGIN {}", zone_origin_substr))
.any(|z| z.origin.contains(zone_origin_substr)); });
if visible == expected_visible { if visible == expected_visible {
return; return;
} }
let origins = snapshot
.zones
.iter()
.map(|z| z.origin.clone())
.collect::<Vec<_>>();
assert!( assert!(
Instant::now() < deadline, Instant::now() < deadline,
"zone visibility mismatch for '{}': expected {}, got {}, current origins: {:?}", "zone visibility mismatch for '{}': expected {}, got {}",
zone_origin_substr, zone_origin_substr,
expected_visible, expected_visible,
visible, visible,
origins
); );
tokio::time::sleep(Duration::from_millis(200)).await; tokio::time::sleep(Duration::from_millis(200)).await;
} }
@@ -504,7 +495,7 @@ records = ["secret IN A 10.99.0.9"]
!snapshot !snapshot
.zones .zones
.iter() .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" "zone without [dns.zone.export] should not be exported to peer snapshot"
); );
} }
@@ -552,7 +543,7 @@ disabled = true
!snapshot !snapshot
.zones .zones
.iter() .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" "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; 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 mut dns = peer.get_global_ctx().config.get_dns().into_raw();
let zone_idx = dns let mut zones = dns.zones.unwrap();
.zones let zone_idx = zones
.iter() .iter()
.position(|z| z.origin.to_string().contains("patch.mesh-test")) .position(|z| z.origin.to_string().contains("patch.mesh-test"))
.expect("patch zone should exist"); .expect("patch zone should exist");
let mut zone: ZoneConfigInner = dns.zones[zone_idx].clone().into(); let mut zone = zones[zone_idx].clone().into_raw();
zone.records = vec!["api IN A 10.80.0.2".to_string()]; zone.records = Some(vec!["api IN A 10.80.0.2".to_string()]);
dns.zones[zone_idx] = zone.try_into().expect("patch zone update should be valid"); zones[zone_idx] = zone.try_into().expect("patch zone update should be valid");
peer.get_global_ctx().config.set_dns(Some(dns)); dns.zones = Some(zones);
peer.get_global_ctx().config.set_dns(dns.into());
peer.get_global_ctx() peer.get_global_ctx()
.issue_event(crate::common::global_ctx::GlobalCtxEvent::ConfigPatched( .issue_event(crate::common::global_ctx::GlobalCtxEvent::ConfigPatched(
crate::proto::api::config::InstanceConfigPatch::default(), 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); 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; 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(); let mut dns = peer.get_global_ctx().config.get_dns().into_raw();
dns.listeners = vec![ dns.listeners = Some(
format!("udp://127.0.0.1:{listener_new}") vec![
.parse() format!("udp://127.0.0.1:{listener_new}")
.expect("invalid listener"), .parse()
] .expect("invalid listener"),
.into(); ]
peer.get_global_ctx().config.set_dns(Some(dns)); .into(),
);
peer.get_global_ctx().config.set_dns(dns.into());
peer.get_global_ctx() peer.get_global_ctx()
.issue_event(crate::common::global_ctx::GlobalCtxEvent::ConfigPatched( .issue_event(crate::common::global_ctx::GlobalCtxEvent::ConfigPatched(
crate::proto::api::config::InstanceConfigPatch::default(), crate::proto::api::config::InstanceConfigPatch::default(),
@@ -894,7 +888,7 @@ records = ["secret IN A 10.66.3.8"]
!snapshot !snapshot
.zones .zones
.iter() .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" "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() .snapshot()
.zones .zones
.iter() .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" "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() .snapshot()
.zones .zones
.iter() .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" "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() .snapshot()
.zones .zones
.iter() .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" "peer B self zone should be restored after DNS RPC rejoins"
); );
} }
+7 -1
View File
@@ -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<NameServerAddr> for Url { impl From<NameServerAddr> for Url {
fn from(value: NameServerAddr) -> Self { fn from(value: NameServerAddr) -> Self {
Url::parse(&format!("{}://{}", value.protocol, value.addr)).unwrap() (&value).into()
} }
} }
+17 -29
View File
@@ -1,6 +1,6 @@
use crate::dns::utils::addr::{NameServerAddr, NameServerAddrGroup}; use crate::dns::utils::addr::{NameServerAddr, NameServerAddrGroup};
use crate::dns::utils::zone_handler::{ArcZoneHandler, ChainedZoneHandler}; use crate::dns::utils::zone_handler::{ArcZoneHandler, ChainedZoneHandler};
use crate::proto; use crate::proto::dns::ZoneData;
use crate::proto::utils::RepeatedMessageModel; use crate::proto::utils::RepeatedMessageModel;
use crate::utils::dns::resolver_conf; use crate::utils::dns::resolver_conf;
use hickory_net::runtime::TokioRuntimeProvider; use hickory_net::runtime::TokioRuntimeProvider;
@@ -13,6 +13,7 @@ use indexmap::IndexMap;
use itertools::chain; use itertools::chain;
use std::collections::BTreeMap; use std::collections::BTreeMap;
use std::sync::Arc; use std::sync::Arc;
use url::Url;
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
pub struct Zone { 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; type Error = anyhow::Error;
fn try_from(value: &proto::dns::ZoneData) -> Result<Self, Self::Error> { fn try_from(value: &ZoneData) -> Result<Self, Self::Error> {
let (origin, records) = Parser::new(value.to_string(), None, None) let (origin, records) = Parser::new(&value.content, None, None)
.parse() .parse()
.map_err(|e| anyhow::anyhow!("failed to parse zone data: {e}"))?; .map_err(|e| anyhow::anyhow!("failed to parse zone data: {e}"))?;
@@ -117,14 +118,13 @@ impl TryFrom<&proto::dns::ZoneData> for Zone {
} }
} }
impl From<Zone> for proto::dns::ZoneData { impl From<Zone> for ZoneData {
fn from(value: Zone) -> Self { fn from(value: Zone) -> Self {
let records = value let records = value
.records .records
.values() .values()
.flat_map(RecordSet::records_without_rrsigs) .flat_map(RecordSet::records_without_rrsigs)
.map(ToString::to_string) .map(ToString::to_string);
.collect();
let forwarders = value let forwarders = value
.forward .forward
@@ -132,16 +132,9 @@ impl From<Zone> for proto::dns::ZoneData {
.flat_map(|f| f.name_servers.into_iter()) .flat_map(|f| f.name_servers.into_iter())
.map(|ns| (&ns).into()) .map(|ns| (&ns).into())
.flat_map(NameServerAddrGroup::into_iter) .flat_map(NameServerAddrGroup::into_iter)
.map(Into::into) .map(Url::from);
.collect();
Self { Self::new(&value.origin, 0, records, forwarders, value.fallthrough)
origin: value.origin.to_string(),
ttl: 0,
records,
forwarders,
fallthrough: value.fallthrough,
}
} }
} }
@@ -202,16 +195,15 @@ mod tests {
forwarders: Vec<&str>, forwarders: Vec<&str>,
fallthrough: bool, fallthrough: bool,
) -> ZoneData { ) -> ZoneData {
ZoneData { ZoneData::new(
origin: origin.to_string(), &origin.parse().unwrap(),
ttl: 60, 60,
records: records.into_iter().map(ToString::to_string).collect(), records,
forwarders: forwarders forwarders.into_iter().map(|url| Url {
.into_iter() url: url.to_string(),
.map(|f| Url::from_str(f).expect("invalid forwarder")) }),
.collect(),
fallthrough, fallthrough,
} )
} }
fn zone_data(origin: &str, records: Vec<&str>, forwarders: Vec<&str>) -> ZoneData { 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); assert_eq!(zone.forward.as_ref().unwrap().name_servers.len(), 2);
let serialized = ZoneData::from(zone.clone()); 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)?; let reparsed = Zone::try_from(&serialized)?;
assert_eq!(reparsed.origin.to_string(), "roundtrip.test."); assert_eq!(reparsed.origin.to_string(), "roundtrip.test.");
assert_eq!(reparsed.iter_records().count(), 2); assert_eq!(reparsed.iter_records().count(), 2);
+3 -5
View File
@@ -5,11 +5,9 @@ import "common.proto";
package dns; package dns;
message ZoneData { message ZoneData {
string origin = 1; string content = 1;
uint32 ttl = 2; repeated common.Url forwarders = 2;
repeated string records = 3; bool fallthrough = 3;
repeated common.Url forwarders = 4;
bool fallthrough = 5;
} }
message GetExportConfigRequest {} message GetExportConfigRequest {}
+38 -23
View File
@@ -1,5 +1,7 @@
use crate::proto::common::Url;
use crate::proto::utils::TransientDigest; 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")); include!(concat!(env!("OUT_DIR"), "/dns.rs"));
@@ -10,31 +12,44 @@ impl HeartbeatRequest {
} }
} }
impl Display for ZoneData { impl ZoneData {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { pub fn new<Records, R, Urls, U>(
writeln!(f, "; EasyTier Magic DNS zone data")?; origin: &LowerName,
writeln!(f, "; https://github.com/easytier/easytier")?; ttl: u32,
records: Records,
forwarders: Urls,
fallthrough: bool,
) -> Self
where
Records: IntoIterator<Item = R>,
R: AsRef<str>,
Urls: IntoIterator<Item = U>,
U: Into<Url>,
{
let mut content = String::new();
if !self.forwarders.is_empty() { content.push_str("; EasyTier Magic DNS zone data\n");
writeln!(f, "; Forwarders:")?; content.push_str("; https://github.com/easytier/easytier\n");
for forwarder in &self.forwarders {
writeln!(f, "; \t{}", forwarder)?;
}
}
writeln!(f)?;
write!(f, "$ORIGIN {}", self.origin)?; let mut origin = origin.to_string();
if !self.origin.ends_with('.') { if !origin.ends_with('.') {
write!(f, ".")?; origin.push('.');
}
writeln!(f)?;
writeln!(f, "$TTL {}", self.ttl)?;
for record in &self.records {
writeln!(f, "{}", record)?;
} }
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,
}
} }
} }
+28 -22
View File
@@ -4127,14 +4127,14 @@ pub async fn three_node_dns_export() {
"tcp", "tcp",
|cfg| { |cfg| {
use crate::dns::config::zone::ZoneConfig; use crate::dns::config::zone::ZoneConfig;
use crate::dns::config::{DnsConfig, DnsConfigLoaderExt}; use crate::dns::config::{DnsConfigLoaderExt, DnsConfigRaw};
use hickory_proto::rr::LowerName; use hickory_proto::rr::LowerName;
use std::str::FromStr; use std::str::FromStr;
let inst_name = cfg.get_inst_name(); let inst_name = cfg.get_inst_name();
let origin = LowerName::from_str(&format!("{}.com.", inst_name)).unwrap(); let origin = LowerName::from_str(&format!("{}.com.", inst_name)).unwrap();
let mut dns_config = DnsConfig { let mut dns_config = DnsConfigRaw {
name: LowerName::from_str(&inst_name).unwrap(), name: Some(LowerName::from_str(&inst_name).unwrap()),
..Default::default() ..Default::default()
}; };
@@ -4147,7 +4147,8 @@ pub async fn three_node_dns_export() {
dns_config dns_config
.zones .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() { let listener_port = match inst_name.as_str() {
"inst1" => 5351, "inst1" => 5351,
@@ -4155,14 +4156,16 @@ pub async fn three_node_dns_export() {
"inst3" => 5353, "inst3" => 5353,
_ => 5350, _ => 5350,
}; };
dns_config.listeners = vec![ dns_config.listeners = Some(
format!("udp://127.0.0.1:{}", listener_port) vec![
.parse() format!("udp://127.0.0.1:{}", listener_port)
.unwrap(), .parse()
] .unwrap(),
.into(); ]
.into(),
);
cfg.set_dns(Some(dns_config)); cfg.set_dns(dns_config.into());
cfg cfg
}, },
false, false,
@@ -4195,14 +4198,14 @@ pub async fn three_node_dns_export_chain() {
let cfg_cb = |cfg: TomlConfigLoader| { let cfg_cb = |cfg: TomlConfigLoader| {
use crate::dns::config::zone::ZoneConfig; use crate::dns::config::zone::ZoneConfig;
use crate::dns::config::{DnsConfig, DnsConfigLoaderExt}; use crate::dns::config::{DnsConfigLoaderExt, DnsConfigRaw};
use hickory_proto::rr::LowerName; use hickory_proto::rr::LowerName;
use std::str::FromStr; use std::str::FromStr;
let inst_name = cfg.get_inst_name(); let inst_name = cfg.get_inst_name();
let origin = LowerName::from_str(&format!("{}.com.", inst_name)).unwrap(); let origin = LowerName::from_str(&format!("{}.com.", inst_name)).unwrap();
let mut dns_config = DnsConfig { let mut dns_config = DnsConfigRaw {
name: LowerName::from_str(&inst_name).unwrap(), name: Some(LowerName::from_str(&inst_name).unwrap()),
..Default::default() ..Default::default()
}; };
@@ -4215,7 +4218,8 @@ pub async fn three_node_dns_export_chain() {
dns_config dns_config
.zones .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() { let listener_port = match inst_name.as_str() {
"inst1" => 5351, "inst1" => 5351,
@@ -4223,14 +4227,16 @@ pub async fn three_node_dns_export_chain() {
"inst3" => 5353, "inst3" => 5353,
_ => 5350, _ => 5350,
}; };
dns_config.listeners = vec![ dns_config.listeners = Some(
format!("udp://127.0.0.1:{}", listener_port) vec![
.parse() format!("udp://127.0.0.1:{}", listener_port)
.unwrap(), .parse()
] .unwrap(),
.into(); ]
.into(),
);
cfg.set_dns(Some(dns_config)); cfg.set_dns(dns_config.into());
let mut flags = cfg.get_flags(); let mut flags = cfg.get_flags();
flags.disable_p2p = true; flags.disable_p2p = true;