From 8a93bb311bd88b982a4509358c0781707125323c Mon Sep 17 00:00:00 2001 From: Luna Yao <40349250+ZnqbuZ@users.noreply.github.com> Date: Sat, 7 Feb 2026 23:12:49 +0100 Subject: [PATCH] add FallbackAuthority and related test --- easytier/src/dns/zone.rs | 229 ++++++++++++++++++++++++++++++++------- 1 file changed, 190 insertions(+), 39 deletions(-) diff --git a/easytier/src/dns/zone.rs b/easytier/src/dns/zone.rs index 5246ba96..86547258 100644 --- a/easytier/src/dns/zone.rs +++ b/easytier/src/dns/zone.rs @@ -1,19 +1,105 @@ use crate::dns::utils::NameServerAddr; use crate::proto::dns::ZoneConfigPb; -use hickory_proto::rr::{LowerName, Record, RecordSet, RrKey, RrsetRecords}; +use async_trait::async_trait; +use derive_more::{Deref, DerefMut}; +use hickory_proto::rr::{LowerName, Record, RecordSet, RecordType, RrKey, RrsetRecords}; use hickory_proto::serialize::txt::Parser; use hickory_proto::xfer::Protocol; use hickory_resolver::config::{NameServerConfig, ResolverOpts}; use hickory_resolver::name_server::TokioConnectionProvider; -use hickory_server::authority::{AuthorityObject, ZoneType}; +use hickory_server::authority::{ + Authority, AuthorityObject, LookupControlFlow, LookupObject, LookupOptions, MessageRequest, + UpdateResult, ZoneType, +}; +use hickory_server::server::RequestInfo; use hickory_server::store::forwarder::{ForwardAuthority, ForwardConfig}; use hickory_server::store::in_memory::InMemoryAuthority; use std::collections::BTreeMap; +use std::mem; use std::net::{IpAddr, SocketAddr}; use std::str::FromStr; use std::sync::Arc; use url::Url; +#[derive(Deref, DerefMut)] +pub struct FallbackAuthority +where + A: Authority + Send + Sync + 'static, + L: LookupObject + Send + Sync + 'static, +{ + #[deref] + #[deref_mut] + inner: A, +} + +#[async_trait] +impl Authority for FallbackAuthority +where + A: Authority + Send + Sync + 'static, + L: LookupObject + Send + Sync + 'static, +{ + type Lookup = L; + + #[inline] + fn zone_type(&self) -> ZoneType { + self.inner.zone_type() + } + #[inline] + fn is_axfr_allowed(&self) -> bool { + self.inner.is_axfr_allowed() + } + #[inline] + async fn update(&self, update: &MessageRequest) -> UpdateResult { + self.inner.update(update).await + } + #[inline] + fn origin(&self) -> &LowerName { + self.inner.origin() + } + #[inline] + async fn lookup( + &self, + name: &LowerName, + rtype: RecordType, + lookup_options: LookupOptions, + ) -> LookupControlFlow { + self.inner.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.inner + .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.inner.search(request_info, lookup_options).await + } + #[inline] + async fn get_nsec_records( + &self, + name: &LowerName, + lookup_options: LookupOptions, + ) -> LookupControlFlow { + self.inner.get_nsec_records(name, lookup_options).await + } +} + #[derive(Debug, Clone)] pub struct Zone { pub(crate) origin: LowerName, @@ -33,13 +119,17 @@ impl Zone { pub fn create_authorities(&self) -> anyhow::Result>> { let mut authorities = Vec::>::with_capacity(2); - let memory = InMemoryAuthority::new( - self.origin.clone().into(), - self.records.clone(), - ZoneType::Primary, - false, - ) - .map_err(|e| anyhow::anyhow!("failed to create memory authority: {}", e))?; + let mut memory = + InMemoryAuthority::empty(self.origin.clone().into(), ZoneType::External, false); + + let mut records = self + .records + .clone() + .into_iter() + .map(|(k, v)| (k, Arc::new(v))) + .collect(); + mem::swap(memory.records_get_mut(), &mut records); + authorities.push(Arc::new(memory)); if let Some(forward) = &self.forward { @@ -49,6 +139,7 @@ impl Zone { ) .build() .map_err(|e| anyhow::anyhow!("failed to create forward authority: {}", e))?; + let forward = FallbackAuthority { inner: forward }; authorities.push(Arc::new(forward)); } @@ -162,21 +253,34 @@ impl TryFrom<&ZoneConfigPb> for Zone { mod tests { use super::*; use crate::dns::config::DnsConfig; - use hickory_proto::rr::{rdata, Name, RData, RecordType}; + use hickory_client::client::{Client, ClientHandle}; + use hickory_proto::rr::{rdata, DNSClass, Name, RData, RecordType}; + use hickory_proto::runtime::TokioRuntimeProvider; + use hickory_proto::udp::UdpClientStream; + use hickory_server::authority::Catalog; + use hickory_server::ServerFuture; + use std::time::Duration; + use tokio::net::UdpSocket; + use tokio::spawn; + use tokio::time::timeout; - #[tokio::test] - async fn config() -> anyhow::Result<()> { - let sep = "=".repeat(80); - let text = r#" + const CONFIG: &str = r#" listeners = [ - "1.1.1.1", + "127.0.0.1:5353", ] name = "et-test" domain = "测试.net" [[zone]] - origin = "et.internal" + origin = "et.top" + + records = [ + "@ 60 A 100.100.100.100", + ] + + [[zone]] + origin = "google.com" broadcast = true @@ -190,25 +294,35 @@ mod tests { ] forwarders = [ - "1.1.1.1", - ] - - [[zone]] - origin = "et.top" - - records = [ - "@ 60 A 100.100.100.100", + "10.175.160.10", ] "#; - let config = toml::from_str::(text)?; + #[tokio::test] + async fn test_config() -> anyhow::Result<()> { + let sep = "=".repeat(80); + let config = toml::from_str::(CONFIG)?; assert_eq!(config.domain.to_string(), "测试.net"); let mut zones = config.zones; assert_eq!(zones.len(), 2); let zone = zones - .extract_if(.., |z| z.origin.to_string() == "et.internal") + .extract_if(.., |c| c.origin.to_string() == "et.top") + .next() + .unwrap(); + let zone = ZoneConfigPb::from(&zone); + let zone = Zone::try_from(&zone)?; + assert_eq!(zone.origin.to_string(), "et.top."); + let records = zone.iter_records().collect::>(); + assert_eq!(records.len(), 1); + + let mut record = Record::update0(zone.origin.clone().into(), 60, RecordType::A); + record.set_data(RData::A(rdata::a::A("100.100.100.100".parse()?))); + assert_eq!(record, **records.iter().next().unwrap()); + + let zone = zones + .extract_if(.., |z| z.origin.to_string() == "google.com") .next() .unwrap(); assert_eq!(zone.broadcast, true); @@ -216,8 +330,11 @@ mod tests { println!("{}", sep); println!("{}", zone); println!("{}", sep); + let zone = Zone::try_from(&zone)?; - assert_eq!(zone.origin.to_string(), "et.internal."); + + assert_eq!(zone.origin.to_string(), "google.com."); + let records = zone.iter_records().collect::>(); assert_eq!(records.len(), 4); @@ -251,19 +368,53 @@ mod tests { .unwrap() ); - let zone = zones - .extract_if(.., |c| c.origin.to_string() == "et.top") - .next() - .unwrap(); - let zone = ZoneConfigPb::from(&zone); - let zone = Zone::try_from(&zone)?; - assert_eq!(zone.origin.to_string(), "et.top."); - let records = zone.iter_records().collect::>(); - assert_eq!(records.len(), 1); + assert_eq!(zone.forward.as_ref().unwrap().name_servers.len(), 1); - let mut record = Record::update0(zone.origin.clone().into(), 60, RecordType::A); - record.set_data(RData::A(rdata::a::A("100.100.100.100".parse()?))); - assert_eq!(record, **records.iter().next().unwrap()); + let authorities = zone.create_authorities()?; + assert_eq!(authorities.len(), 2); + let mut catalog = Catalog::new(); + catalog.upsert(zone.origin.clone().into(), authorities); + + let socket = UdpSocket::bind("127.0.0.1:0").await?; + let addr = socket.local_addr()?; + println!("listening on {}", addr); + + let mut server = ServerFuture::new(catalog); + server.register_socket(socket); + spawn(async move { + if let Err(e) = server.block_until_done().await { + eprintln!("server error: {}", e); + } + }); + + let conn = UdpClientStream::builder(addr, TokioRuntimeProvider::default()).build(); + let (mut client, background) = + timeout(Duration::from_secs(1), Client::connect(conn)).await??; + spawn(async move { + if let Err(e) = background.await { + eprintln!("client error: {}", e); + } + }); + let value = timeout( + Duration::from_secs(1), + client.query("maps.google.com".parse()?, DNSClass::IN, RecordType::A), + ) + .await??; + value.answers().iter().for_each(|r| println!("{}", r)); + + let value = timeout( + Duration::from_secs(1), + client.query("www.google.com".parse()?, DNSClass::IN, RecordType::A), + ) + .await??; + value.answers().iter().for_each(|r| println!("{}", r)); + + let value = timeout( + Duration::from_secs(1), + client.query("google.com".parse()?, DNSClass::IN, RecordType::A), + ) + .await??; + value.answers().iter().for_each(|r| println!("{}", r)); Ok(()) }