add FallbackAuthority and related test

This commit is contained in:
Luna Yao
2026-04-06 11:54:01 +02:00
parent b24bb42faa
commit 8a93bb311b
+190 -39
View File
@@ -1,19 +1,105 @@
use crate::dns::utils::NameServerAddr; use crate::dns::utils::NameServerAddr;
use crate::proto::dns::ZoneConfigPb; 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::serialize::txt::Parser;
use hickory_proto::xfer::Protocol; use hickory_proto::xfer::Protocol;
use hickory_resolver::config::{NameServerConfig, ResolverOpts}; use hickory_resolver::config::{NameServerConfig, ResolverOpts};
use hickory_resolver::name_server::TokioConnectionProvider; 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::forwarder::{ForwardAuthority, ForwardConfig};
use hickory_server::store::in_memory::InMemoryAuthority; use hickory_server::store::in_memory::InMemoryAuthority;
use std::collections::BTreeMap; use std::collections::BTreeMap;
use std::mem;
use std::net::{IpAddr, SocketAddr}; use std::net::{IpAddr, SocketAddr};
use std::str::FromStr; use std::str::FromStr;
use std::sync::Arc; use std::sync::Arc;
use url::Url; use url::Url;
#[derive(Deref, DerefMut)]
pub struct FallbackAuthority<A, L>
where
A: Authority<Lookup = L> + Send + Sync + 'static,
L: LookupObject + Send + Sync + 'static,
{
#[deref]
#[deref_mut]
inner: A,
}
#[async_trait]
impl<A, L> Authority for FallbackAuthority<A, L>
where
A: Authority<Lookup = L> + 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<bool> {
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::Lookup> {
self.inner.lookup(name, rtype, lookup_options).await
}
#[inline]
async fn consult(
&self,
name: &LowerName,
rtype: RecordType,
lookup_options: LookupOptions,
last_result: LookupControlFlow<Box<dyn LookupObject>>,
) -> LookupControlFlow<Box<dyn LookupObject>> {
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::Lookup> {
self.inner.search(request_info, lookup_options).await
}
#[inline]
async fn get_nsec_records(
&self,
name: &LowerName,
lookup_options: LookupOptions,
) -> LookupControlFlow<Self::Lookup> {
self.inner.get_nsec_records(name, lookup_options).await
}
}
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
pub struct Zone { pub struct Zone {
pub(crate) origin: LowerName, pub(crate) origin: LowerName,
@@ -33,13 +119,17 @@ impl Zone {
pub fn create_authorities(&self) -> anyhow::Result<Vec<Arc<dyn AuthorityObject>>> { pub fn create_authorities(&self) -> anyhow::Result<Vec<Arc<dyn AuthorityObject>>> {
let mut authorities = Vec::<Arc<dyn AuthorityObject>>::with_capacity(2); let mut authorities = Vec::<Arc<dyn AuthorityObject>>::with_capacity(2);
let memory = InMemoryAuthority::new( let mut memory =
self.origin.clone().into(), InMemoryAuthority::empty(self.origin.clone().into(), ZoneType::External, false);
self.records.clone(),
ZoneType::Primary, let mut records = self
false, .records
) .clone()
.map_err(|e| anyhow::anyhow!("failed to create memory authority: {}", e))?; .into_iter()
.map(|(k, v)| (k, Arc::new(v)))
.collect();
mem::swap(memory.records_get_mut(), &mut records);
authorities.push(Arc::new(memory)); authorities.push(Arc::new(memory));
if let Some(forward) = &self.forward { if let Some(forward) = &self.forward {
@@ -49,6 +139,7 @@ impl Zone {
) )
.build() .build()
.map_err(|e| anyhow::anyhow!("failed to create forward authority: {}", e))?; .map_err(|e| anyhow::anyhow!("failed to create forward authority: {}", e))?;
let forward = FallbackAuthority { inner: forward };
authorities.push(Arc::new(forward)); authorities.push(Arc::new(forward));
} }
@@ -162,21 +253,34 @@ impl TryFrom<&ZoneConfigPb> for Zone {
mod tests { mod tests {
use super::*; use super::*;
use crate::dns::config::DnsConfig; 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] const CONFIG: &str = r#"
async fn config() -> anyhow::Result<()> {
let sep = "=".repeat(80);
let text = r#"
listeners = [ listeners = [
"1.1.1.1", "127.0.0.1:5353",
] ]
name = "et-test" name = "et-test"
domain = "测试.net" domain = "测试.net"
[[zone]] [[zone]]
origin = "et.internal" origin = "et.top"
records = [
"@ 60 A 100.100.100.100",
]
[[zone]]
origin = "google.com"
broadcast = true broadcast = true
@@ -190,25 +294,35 @@ mod tests {
] ]
forwarders = [ forwarders = [
"1.1.1.1", "10.175.160.10",
]
[[zone]]
origin = "et.top"
records = [
"@ 60 A 100.100.100.100",
] ]
"#; "#;
let config = toml::from_str::<DnsConfig>(text)?; #[tokio::test]
async fn test_config() -> anyhow::Result<()> {
let sep = "=".repeat(80);
let config = toml::from_str::<DnsConfig>(CONFIG)?;
assert_eq!(config.domain.to_string(), "测试.net"); assert_eq!(config.domain.to_string(), "测试.net");
let mut zones = config.zones; let mut zones = config.zones;
assert_eq!(zones.len(), 2); assert_eq!(zones.len(), 2);
let zone = zones 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::<Vec<_>>();
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() .next()
.unwrap(); .unwrap();
assert_eq!(zone.broadcast, true); assert_eq!(zone.broadcast, true);
@@ -216,8 +330,11 @@ mod tests {
println!("{}", sep); println!("{}", sep);
println!("{}", zone); println!("{}", zone);
println!("{}", sep); println!("{}", sep);
let zone = Zone::try_from(&zone)?; 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::<Vec<_>>(); let records = zone.iter_records().collect::<Vec<_>>();
assert_eq!(records.len(), 4); assert_eq!(records.len(), 4);
@@ -251,19 +368,53 @@ mod tests {
.unwrap() .unwrap()
); );
let zone = zones assert_eq!(zone.forward.as_ref().unwrap().name_servers.len(), 1);
.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::<Vec<_>>();
assert_eq!(records.len(), 1);
let mut record = Record::update0(zone.origin.clone().into(), 60, RecordType::A); let authorities = zone.create_authorities()?;
record.set_data(RData::A(rdata::a::A("100.100.100.100".parse()?))); assert_eq!(authorities.len(), 2);
assert_eq!(record, **records.iter().next().unwrap()); 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(()) Ok(())
} }