diff --git a/easytier/src/dns/zone.rs b/easytier/src/dns/zone.rs index 66e087df..7dcf365c 100644 --- a/easytier/src/dns/zone.rs +++ b/easytier/src/dns/zone.rs @@ -1,7 +1,9 @@ use crate::common::dns::get_default_resolver_config; use crate::dns::utils::NameServerAddr; use crate::proto; +use crate::proto::utils::RepeatedMessageModel; use crate::utils::MapTryInto; +use derivative::Derivative; use hickory_proto::rr::{LowerName, Record, RecordSet, RrKey, RrsetRecords}; use hickory_proto::serialize::txt::Parser; use hickory_resolver::config::ResolverOpts; @@ -15,11 +17,13 @@ use std::collections::BTreeMap; use std::sync::Arc; use uuid::Uuid; -#[derive(Debug, Clone)] +#[derive(Derivative, Debug, Clone)] +#[derivative(PartialEq)] pub struct Zone { pub(crate) id: Uuid, pub(crate) origin: LowerName, pub(crate) records: BTreeMap, + #[derivative(PartialEq(compare_with = "Zone::compare_forward"))] pub(crate) forward: Option, } @@ -35,6 +39,19 @@ impl Zone { zone.forward = Some(forward); zone } + + pub fn compare_forward(l: &Option, r: &Option) -> bool { + match (l, r) { + (Some(l), Some(r)) => l + .name_servers + .iter() + .cloned() + .map_into::() + .eq(r.name_servers.iter().cloned().map_into()), + (None, None) => true, + _ => false, + } + } } impl Zone { @@ -108,7 +125,7 @@ impl TryFrom<&proto::dns::ZoneData> for Zone { .iter() .map_try_into::() .map_ok(Into::into) - .collect::>>()? + .try_collect::<_, Vec<_>, _>()? .into(); let forward = Some(ForwardConfig { name_servers, @@ -150,6 +167,8 @@ impl From for proto::dns::ZoneData { } } +pub type ZoneGroup = RepeatedMessageModel; + #[cfg(test)] mod tests { use super::*; @@ -222,7 +241,7 @@ mod tests { let zone = proto::dns::ZoneData::from(zone); let zone = Zone::try_from(&zone)?; assert_eq!(zone.origin.to_string(), "et.top."); - let records = zone.iter_records().collect::>(); + let records = zone.iter_records().collect_vec(); assert_eq!(records.len(), 1); let mut record = Record::update0(zone.origin.clone().into(), 60, RecordType::A); @@ -247,7 +266,7 @@ mod tests { assert_eq!(zone.origin.to_string(), "google.com."); - let records = zone.iter_records().collect::>(); + let records = zone.iter_records().collect_vec(); assert_eq!(records.len(), 4); let mut record = Record::update0(