use std::collections::HashMap;
use std::net::{Ipv4Addr, Ipv6Addr};
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex};
use std::time::{Duration, Instant};
use dns_lattice_core::{Error, Result};
use dns_lattice_model::{
Class, Message, Name, RData, Rcode, RecordType, ResourceRecord, SplitDnsPolicy, UpstreamGroupId,
};
use crate::fakeip::{FakeIpPolicy, FakeIpPool};
use crate::hooks::{RouteDecision, RouteHook, RouteRequest};
use crate::observability::{
HookObserveDecision, ObservabilitySink, ObserveEvent, ObserveFailure, UpstreamObserveOutcome,
};
use crate::upstream::UpstreamBackend;
const NEGATIVE_CACHE_FLOOR: Duration = Duration::from_secs(60);
pub(crate) trait Clock {
fn now(&self) -> Instant;
}
pub(crate) struct SystemClock;
impl Clock for SystemClock {
fn now(&self) -> Instant {
Instant::now()
}
}
#[cfg(test)]
#[derive(Clone)]
pub(crate) struct FakeClock(std::sync::Arc<std::sync::Mutex<Instant>>);
#[cfg(test)]
impl FakeClock {
pub(crate) fn new() -> Self {
FakeClock(std::sync::Arc::new(std::sync::Mutex::new(Instant::now())))
}
pub(crate) fn advance(&self, duration: Duration) {
let mut guard = self.0.lock().expect("fake clock mutex poisoned");
*guard += duration;
}
}
#[cfg(test)]
impl Clock for FakeClock {
fn now(&self) -> Instant {
*self.0.lock().expect("fake clock mutex poisoned")
}
}
#[derive(Clone, PartialEq, Eq, Hash)]
struct CacheKey {
name: Name,
rtype: RecordType,
class: Class,
group: UpstreamGroupId,
}
struct CacheEntry {
answer: Message,
expires_at: Instant,
}
pub struct Resolver {
policy: SplitDnsPolicy,
backends: HashMap<UpstreamGroupId, Vec<Box<dyn UpstreamBackend>>>,
clock: Box<dyn Clock + Send + Sync>,
cache: Mutex<HashMap<CacheKey, CacheEntry>>,
fake_ip: Option<FakeIpResolverConfig>,
route_hook: Option<Box<dyn RouteHook>>,
observability_sink: Option<Arc<dyn ObservabilitySink>>,
next_correlation_id: AtomicU64,
}
struct FakeIpResolverConfig {
pool: Arc<FakeIpPool>,
policy: FakeIpPolicy,
}
impl Resolver {
pub fn builder(policy: SplitDnsPolicy) -> ResolverBuilder {
ResolverBuilder {
policy,
backends: HashMap::new(),
clock: Box::new(SystemClock),
fake_ip: None,
route_hook: None,
observability_sink: None,
}
}
pub async fn resolve(&self, query: &Message) -> Result<Message> {
let correlation_id = self.next_correlation_id.fetch_add(1, Ordering::Relaxed);
let question = query.questions.first();
self.emit(ObserveEvent::QueryReceived {
correlation_id,
name: question.map(|question| question.name.clone()),
rtype: question.map(|question| question.qtype),
class: question.map(|question| question.qclass),
});
let Some(question) = query.questions.first() else {
self.emit(ObserveEvent::Failed {
correlation_id,
failure: ObserveFailure::NoRoute,
});
return Err(Error::NoRoute);
};
if let Some(fake_ip) = &self.fake_ip {
match fake_ip_answer(query, fake_ip) {
Ok(Some(answer)) => {
self.emit(ObserveEvent::FakeIpTerminal { correlation_id });
self.emit(ObserveEvent::Completed {
correlation_id,
rcode: answer.header.rcode,
});
return Ok(answer);
}
Ok(None) => {}
Err(error) => {
self.emit(ObserveEvent::Failed {
correlation_id,
failure: observe_failure(&error),
});
return Err(error);
}
}
}
let (group, backends) = match self.select_backends(question, correlation_id).await {
Ok(selected) => selected,
Err(error) => {
self.emit(ObserveEvent::Failed {
correlation_id,
failure: observe_failure(&error),
});
return Err(error);
}
};
let key = CacheKey {
name: question.name.clone(),
rtype: question.qtype,
class: question.qclass,
group: group.clone(),
};
let now = self.clock.now();
{
let cache = self.cache.lock().expect("cache mutex poisoned");
if let Some(entry) = cache.get(&key)
&& entry.expires_at > now
{
let answer = cache_hit_response(query, &entry.answer);
drop(cache);
self.emit(ObserveEvent::CacheHit {
correlation_id,
group: group.clone(),
});
self.emit(ObserveEvent::Completed {
correlation_id,
rcode: answer.header.rcode,
});
return Ok(answer);
}
}
self.emit(ObserveEvent::CacheMiss {
correlation_id,
group: group.clone(),
});
let mut last_err = None;
for (backend_index, backend) in backends.iter().enumerate() {
self.emit(ObserveEvent::UpstreamAttempt {
correlation_id,
group: group.clone(),
backend_index,
});
match backend.resolve(query).await {
Ok(answer) => {
self.emit(ObserveEvent::UpstreamOutcome {
correlation_id,
group: group.clone(),
backend_index,
outcome: UpstreamObserveOutcome::Success,
});
if let Some(ttl) = cacheable_ttl(&answer) {
let mut cache = self.cache.lock().expect("cache mutex poisoned");
cache.insert(
key,
CacheEntry {
answer: answer.clone(),
expires_at: now + ttl,
},
);
}
self.emit(ObserveEvent::Completed {
correlation_id,
rcode: answer.header.rcode,
});
return Ok(answer);
}
Err(e) if is_retryable(&e) => {
self.emit(ObserveEvent::UpstreamOutcome {
correlation_id,
group: group.clone(),
backend_index,
outcome: UpstreamObserveOutcome::RetryableFailure,
});
last_err = Some(e);
}
Err(e) => {
self.emit(ObserveEvent::UpstreamOutcome {
correlation_id,
group: group.clone(),
backend_index,
outcome: UpstreamObserveOutcome::Failure,
});
self.emit(ObserveEvent::Failed {
correlation_id,
failure: observe_failure(&e),
});
return Err(e);
}
}
}
let error = last_err.expect("at least one backend was tried since backends is non-empty");
self.emit(ObserveEvent::Failed {
correlation_id,
failure: observe_failure(&error),
});
Err(error)
}
fn emit(&self, event: ObserveEvent) {
if let Some(sink) = &self.observability_sink {
let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| sink.record(&event)));
}
}
async fn select_backends(
&self,
question: &dns_lattice_model::Question,
correlation_id: u64,
) -> Result<(UpstreamGroupId, &Vec<Box<dyn UpstreamBackend>>)> {
let static_group = self.policy.resolve_group(&question.name);
self.emit(ObserveEvent::StaticRoute {
correlation_id,
group: static_group.cloned(),
});
let group = match &self.route_hook {
Some(hook) => match hook.select(RouteRequest::new(question, static_group)).await {
Ok(RouteDecision::Use(group)) => {
self.emit(ObserveEvent::HookDecision {
correlation_id,
decision: HookObserveDecision::Use(group.clone()),
});
Some(group)
}
Ok(RouteDecision::Abstain) => {
self.emit(ObserveEvent::HookDecision {
correlation_id,
decision: HookObserveDecision::Abstain,
});
static_group.cloned()
}
Err(error) => {
self.emit(ObserveEvent::HookDecision {
correlation_id,
decision: HookObserveDecision::Failed,
});
return Err(Error::Hook(error.to_string()));
}
},
None => static_group.cloned(),
}
.ok_or(Error::NoRoute)?;
let backends = self.backends.get(&group).ok_or(Error::NoRoute)?;
if backends.is_empty() {
return Err(Error::NoRoute);
}
Ok((group, backends))
}
}
fn cache_hit_response(query: &Message, cached: &Message) -> Message {
let mut response = cached.clone();
response.header.id = query.header.id;
response.questions = query.questions.clone();
response
}
fn fake_ip_answer(query: &Message, fake_ip: &FakeIpResolverConfig) -> Result<Option<Message>> {
let Some(question) = query.questions.first() else {
return Ok(None);
};
if question.qclass != Class::In {
return Ok(None);
}
match question.qtype {
RecordType::A if fake_ip.policy.matches(&question.name) => {
if !fake_ip.pool.ipv4_enabled() {
return Ok(Some(local_response(query, Rcode::NoError)));
}
fake_ip_ttl(fake_ip.pool.ttl())?;
let mut answer = local_response(query, Rcode::NoError);
match fake_ip.pool.allocate_ipv4_with_ttl(question.name.clone()) {
Ok((address, lifetime)) => answer.answers.push(ResourceRecord {
name: question.name.clone(),
rtype: RecordType::A,
class: Class::In,
ttl: fake_ip_ttl(lifetime)?,
rdata: RData::A(address),
}),
Err(Error::FakeIpFamilyDisabled) => {}
Err(error) => return Err(error),
}
Ok(Some(answer))
}
RecordType::Aaaa if fake_ip.policy.matches(&question.name) => {
if !fake_ip.pool.ipv6_enabled() {
return Ok(Some(local_response(query, Rcode::NoError)));
}
fake_ip_ttl(fake_ip.pool.ttl())?;
let mut answer = local_response(query, Rcode::NoError);
match fake_ip.pool.allocate_ipv6_with_ttl(question.name.clone()) {
Ok((address, lifetime)) => answer.answers.push(ResourceRecord {
name: question.name.clone(),
rtype: RecordType::Aaaa,
class: Class::In,
ttl: fake_ip_ttl(lifetime)?,
rdata: RData::Aaaa(address),
}),
Err(Error::FakeIpFamilyDisabled) => {}
Err(error) => return Err(error),
}
Ok(Some(answer))
}
RecordType::Ptr => fake_ip_ptr_answer(query, fake_ip),
_ => Ok(None),
}
}
fn fake_ip_ptr_answer(query: &Message, fake_ip: &FakeIpResolverConfig) -> Result<Option<Message>> {
let question = query.questions.first().expect("checked by caller");
let address = match parse_reverse_name(&question.name) {
Some(address) => address,
None => return Ok(None),
};
let mapping = match address {
std::net::IpAddr::V4(address) if fake_ip.pool.contains_ipv4(address) => {
fake_ip.pool.lookup_ipv4_with_ttl(address)
}
std::net::IpAddr::V6(address) if fake_ip.pool.contains_ipv6(address) => {
fake_ip.pool.lookup_ipv6_with_ttl(address)
}
_ => return Ok(None),
};
let mut answer = local_response(
query,
if mapping.is_some() {
Rcode::NoError
} else {
Rcode::NxDomain
},
);
if let Some((name, lifetime)) = mapping {
answer.answers.push(ResourceRecord {
name: question.name.clone(),
rtype: RecordType::Ptr,
class: Class::In,
ttl: fake_ip_ttl(lifetime)?,
rdata: RData::Ptr(name),
});
}
Ok(Some(answer))
}
fn local_response(query: &Message, rcode: Rcode) -> Message {
let mut header = query.header;
header.qr = true;
header.rcode = rcode;
Message {
header,
questions: query.questions.clone(),
answers: Vec::new(),
authorities: Vec::new(),
additionals: Vec::new(),
}
}
fn fake_ip_ttl(lifetime: Duration) -> Result<u32> {
u32::try_from(lifetime.as_secs()).map_err(|_| Error::FakeIpTtlOutOfRange)
}
fn parse_reverse_name(name: &Name) -> Option<std::net::IpAddr> {
let labels: Vec<_> = name.labels().collect();
if labels.len() == 6
&& labels[4].eq_ignore_ascii_case(b"in-addr")
&& labels[5].eq_ignore_ascii_case(b"arpa")
{
let mut octets = [0_u8; 4];
for (index, label) in labels[..4].iter().enumerate() {
let text = std::str::from_utf8(label).ok()?;
let value = text.parse::<u8>().ok()?;
if value.to_string() != text {
return None;
}
octets[3 - index] = value;
}
return Some(std::net::IpAddr::V4(Ipv4Addr::from(octets)));
}
if labels.len() == 34
&& labels[32].eq_ignore_ascii_case(b"ip6")
&& labels[33].eq_ignore_ascii_case(b"arpa")
{
let mut bytes = [0_u8; 16];
for (index, label) in labels[..32].iter().enumerate() {
if label.len() != 1 {
return None;
}
let nibble = match label[0] {
b'0'..=b'9' => label[0] - b'0',
b'a'..=b'f' => label[0] - b'a' + 10,
b'A'..=b'F' => label[0] - b'A' + 10,
_ => return None,
};
let target = 31 - index;
if target % 2 == 0 {
bytes[target / 2] |= nibble << 4;
} else {
bytes[target / 2] |= nibble;
}
}
return Some(std::net::IpAddr::V6(Ipv6Addr::from(bytes)));
}
None
}
fn is_retryable(err: &Error) -> bool {
matches!(err, Error::Timeout | Error::Transport(_) | Error::Tls(_))
}
fn observe_failure(error: &Error) -> ObserveFailure {
match error {
Error::NoRoute => ObserveFailure::NoRoute,
Error::Hook(_) => ObserveFailure::Hook,
Error::Timeout => ObserveFailure::Timeout,
Error::Transport(_) => ObserveFailure::Transport,
Error::Tls(_) => ObserveFailure::Tls,
_ => ObserveFailure::Other,
}
}
fn cacheable_ttl(answer: &Message) -> Option<Duration> {
let is_negative = matches!(answer.header.rcode, Rcode::NxDomain)
|| (matches!(answer.header.rcode, Rcode::NoError) && answer.answers.is_empty());
if is_negative {
let ttl = answer
.authorities
.iter()
.find_map(|rr| match &rr.rdata {
RData::Soa { minimum, .. } => Some(*minimum),
_ => None,
})
.map(|minimum| Duration::from_secs(u64::from(minimum)))
.unwrap_or(NEGATIVE_CACHE_FLOOR);
return Some(ttl);
}
if answer.answers.is_empty() {
return None;
}
answer
.answers
.iter()
.map(|rr| rr.ttl)
.min()
.map(|ttl| Duration::from_secs(u64::from(ttl)))
}
pub struct ResolverBuilder {
policy: SplitDnsPolicy,
backends: HashMap<UpstreamGroupId, Vec<Box<dyn UpstreamBackend>>>,
clock: Box<dyn Clock + Send + Sync>,
fake_ip: Option<FakeIpResolverConfig>,
route_hook: Option<Box<dyn RouteHook>>,
observability_sink: Option<Arc<dyn ObservabilitySink>>,
}
impl ResolverBuilder {
pub fn backend(
mut self,
group: UpstreamGroupId,
backend: impl UpstreamBackend + 'static,
) -> Self {
self.backends
.entry(group)
.or_default()
.push(Box::new(backend));
self
}
pub fn fake_ip(mut self, pool: Arc<FakeIpPool>, policy: FakeIpPolicy) -> Self {
self.fake_ip = Some(FakeIpResolverConfig { pool, policy });
self
}
pub fn route_hook(mut self, hook: impl RouteHook + 'static) -> Self {
self.route_hook = Some(Box::new(hook));
self
}
pub fn observability_sink(mut self, sink: Arc<dyn ObservabilitySink>) -> Self {
self.observability_sink = Some(sink);
self
}
#[cfg(test)]
pub(crate) fn clock(mut self, clock: impl Clock + Send + Sync + 'static) -> Self {
self.clock = Box::new(clock);
self
}
pub fn build(self) -> Resolver {
Resolver {
policy: self.policy,
backends: self.backends,
clock: self.clock,
cache: Mutex::new(HashMap::new()),
fake_ip: self.fake_ip,
route_hook: self.route_hook,
observability_sink: self.observability_sink,
next_correlation_id: AtomicU64::new(1),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use async_trait::async_trait;
use dns_lattice_model::{
Class, DomainPattern, Header, Name, Opcode, Question, Rcode, RecordType,
};
struct FixedBackend(Message);
#[async_trait]
impl UpstreamBackend for FixedBackend {
async fn resolve(&self, _query: &Message) -> Result<Message> {
Ok(self.0.clone())
}
}
fn fixed_backend(answer: Message) -> FixedBackend {
FixedBackend(answer)
}
struct FailingBackend(Error);
#[async_trait]
impl UpstreamBackend for FailingBackend {
async fn resolve(&self, _query: &Message) -> Result<Message> {
Err(self.0.clone())
}
}
fn n(s: &str) -> Name {
Name::from_ascii(s).unwrap()
}
fn query_for(name: &str) -> Message {
Message {
header: Header {
id: 1,
qr: false,
opcode: Opcode::Query,
authoritative: false,
truncated: false,
recursion_desired: true,
recursion_available: false,
rcode: Rcode::NoError,
},
questions: vec![Question {
name: n(name),
qtype: RecordType::A,
qclass: Class::In,
}],
answers: vec![],
authorities: vec![],
additionals: vec![],
}
}
fn answer_tagged(id: u16) -> Message {
let mut msg = query_for("tag.example");
msg.header.id = id;
msg.header.qr = true;
msg
}
#[tokio::test]
async fn routes_exact_match_to_its_group() {
let policy = SplitDnsPolicy::builder()
.rule(
DomainPattern::exact(n("host.corp.internal")),
UpstreamGroupId::new("corp"),
)
.build();
let resolver = Resolver::builder(policy)
.backend(
UpstreamGroupId::new("corp"),
fixed_backend(answer_tagged(42)),
)
.build();
let answer = resolver
.resolve(&query_for("host.corp.internal"))
.await
.expect("routed to corp backend");
assert_eq!(answer.header.id, 42);
}
#[tokio::test]
async fn routes_suffix_match_to_its_group() {
let policy = SplitDnsPolicy::builder()
.rule(
DomainPattern::suffix(n("corp.internal")),
UpstreamGroupId::new("corp"),
)
.build();
let resolver = Resolver::builder(policy)
.backend(
UpstreamGroupId::new("corp"),
fixed_backend(answer_tagged(7)),
)
.build();
let answer = resolver
.resolve(&query_for("host.corp.internal"))
.await
.expect("routed to corp backend via suffix");
assert_eq!(answer.header.id, 7);
}
#[tokio::test]
async fn routes_wildcard_match_to_its_group() {
let policy = SplitDnsPolicy::builder()
.rule(
DomainPattern::wildcard(n("corp.internal")),
UpstreamGroupId::new("corp"),
)
.build();
let resolver = Resolver::builder(policy)
.backend(
UpstreamGroupId::new("corp"),
fixed_backend(answer_tagged(9)),
)
.build();
let answer = resolver
.resolve(&query_for("host.corp.internal"))
.await
.expect("routed to corp backend via wildcard");
assert_eq!(answer.header.id, 9);
}
#[tokio::test]
async fn routes_unmatched_query_to_default_group() {
let policy = SplitDnsPolicy::builder()
.rule(
DomainPattern::suffix(n("corp.internal")),
UpstreamGroupId::new("corp"),
)
.default_group(UpstreamGroupId::new("public"))
.build();
let resolver = Resolver::builder(policy)
.backend(
UpstreamGroupId::new("public"),
fixed_backend(answer_tagged(3)),
)
.build();
let answer = resolver
.resolve(&query_for("example.com"))
.await
.expect("routed to default group");
assert_eq!(answer.header.id, 3);
}
#[tokio::test]
async fn no_route_when_no_match_and_no_default_group() {
let policy = SplitDnsPolicy::builder().build();
let resolver = Resolver::builder(policy).build();
let err = resolver
.resolve(&query_for("example.com"))
.await
.expect_err("no rule and no default group configured");
assert_eq!(err, Error::NoRoute);
}
#[tokio::test]
async fn no_route_when_matched_group_has_no_registered_backend() {
let policy = SplitDnsPolicy::builder()
.rule(
DomainPattern::suffix(n("corp.internal")),
UpstreamGroupId::new("corp"),
)
.build();
let resolver = Resolver::builder(policy).build();
let err = resolver
.resolve(&query_for("host.corp.internal"))
.await
.expect_err("matched group has no backend registered");
assert_eq!(err, Error::NoRoute);
}
#[tokio::test]
async fn failover_first_backend_succeeds_second_never_called() {
let policy = SplitDnsPolicy::builder()
.default_group(UpstreamGroupId::new("g"))
.build();
let calls = Arc::new(AtomicUsize::new(0));
let second = CountingBackend {
answer: answer_tagged(2),
calls: calls.clone(),
};
let resolver = Resolver::builder(policy)
.backend(UpstreamGroupId::new("g"), fixed_backend(answer_tagged(1)))
.backend(UpstreamGroupId::new("g"), second)
.build();
let answer = resolver
.resolve(&query_for("example.com"))
.await
.expect("first backend answers");
assert_eq!(answer.header.id, 1);
assert_eq!(
calls.load(Ordering::SeqCst),
0,
"second backend never called once the first succeeds"
);
}
#[tokio::test]
async fn failover_first_backend_fails_second_succeeds() {
let policy = SplitDnsPolicy::builder()
.default_group(UpstreamGroupId::new("g"))
.build();
let resolver = Resolver::builder(policy)
.backend(UpstreamGroupId::new("g"), FailingBackend(Error::Timeout))
.backend(UpstreamGroupId::new("g"), fixed_backend(answer_tagged(99)))
.build();
let answer = resolver
.resolve(&query_for("example.com"))
.await
.expect("second backend answers after first times out");
assert_eq!(
answer.header.id, 99,
"routed answer is the second backend's"
);
}
#[tokio::test]
async fn failover_tls_error_retries_to_next_backend() {
let policy = SplitDnsPolicy::builder()
.default_group(UpstreamGroupId::new("g"))
.build();
let resolver = Resolver::builder(policy)
.backend(
UpstreamGroupId::new("g"),
FailingBackend(Error::Tls("certificate expired".to_string())),
)
.backend(UpstreamGroupId::new("g"), fixed_backend(answer_tagged(5)))
.build();
let answer = resolver
.resolve(&query_for("example.com"))
.await
.expect("tls error on first backend retries to the second");
assert_eq!(answer.header.id, 5);
}
#[tokio::test]
async fn failover_all_backends_fail_returns_last_error() {
let policy = SplitDnsPolicy::builder()
.default_group(UpstreamGroupId::new("g"))
.build();
let resolver = Resolver::builder(policy)
.backend(UpstreamGroupId::new("g"), FailingBackend(Error::Timeout))
.backend(
UpstreamGroupId::new("g"),
FailingBackend(Error::Transport("connection refused".to_string())),
)
.build();
let err = resolver
.resolve(&query_for("example.com"))
.await
.expect_err("both backends fail");
assert_eq!(
err,
Error::Transport("connection refused".to_string()),
"the last attempted backend's error is returned, not the first's"
);
}
#[tokio::test]
async fn single_backend_group_still_behaves_as_before() {
let policy = SplitDnsPolicy::builder()
.default_group(UpstreamGroupId::new("g"))
.build();
let resolver = Resolver::builder(policy)
.backend(UpstreamGroupId::new("g"), fixed_backend(answer_tagged(11)))
.build();
let answer = resolver
.resolve(&query_for("example.com"))
.await
.expect("single-backend group still resolves");
assert_eq!(answer.header.id, 11);
}
#[tokio::test]
async fn backend_error_propagates_as_is() {
let policy = SplitDnsPolicy::builder()
.rule(
DomainPattern::suffix(n("corp.internal")),
UpstreamGroupId::new("corp"),
)
.build();
let resolver = Resolver::builder(policy)
.backend(
UpstreamGroupId::new("corp"),
FailingBackend(Error::NameTooLong),
)
.build();
let err = resolver
.resolve(&query_for("host.corp.internal"))
.await
.expect_err("backend failure propagates");
assert_eq!(err, Error::NameTooLong);
}
use std::net::Ipv4Addr;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use dns_lattice_model::{RData, ResourceRecord};
use tokio::sync::Notify;
use crate::hooks::{RouteDecision, RouteHook, RouteHookError, RouteRequest};
use crate::observability::{ObservabilitySink, ObserveEvent};
#[derive(Clone)]
struct PoolClock(Arc<Mutex<Instant>>);
impl PoolClock {
fn new() -> Self {
Self(Arc::new(Mutex::new(Instant::now())))
}
fn advance(&self, duration: Duration) {
*self.0.lock().expect("pool clock mutex poisoned") += duration;
}
}
impl crate::fakeip::Clock for PoolClock {
fn now(&self) -> Instant {
*self.0.lock().expect("pool clock mutex poisoned")
}
}
struct CountingBackend {
answer: Message,
calls: Arc<AtomicUsize>,
}
#[derive(Default)]
struct RecordingSink(Mutex<Vec<ObserveEvent>>);
impl ObservabilitySink for RecordingSink {
fn record(&self, event: &ObserveEvent) {
self.0
.lock()
.expect("event mutex poisoned")
.push(event.clone());
}
}
struct PanickingSink;
impl ObservabilitySink for PanickingSink {
fn record(&self, _: &ObserveEvent) {
panic!("observer failure must be isolated");
}
}
#[async_trait]
impl UpstreamBackend for CountingBackend {
async fn resolve(&self, _query: &Message) -> Result<Message> {
self.calls.fetch_add(1, Ordering::SeqCst);
Ok(self.answer.clone())
}
}
struct FixedHook {
decision: std::result::Result<RouteDecision, RouteHookError>,
calls: Arc<AtomicUsize>,
}
#[async_trait]
impl RouteHook for FixedHook {
async fn select(
&self,
_request: RouteRequest<'_>,
) -> std::result::Result<RouteDecision, RouteHookError> {
self.calls.fetch_add(1, Ordering::SeqCst);
self.decision.clone()
}
}
struct SequencedHook {
decisions: Mutex<Vec<RouteDecision>>,
}
#[async_trait]
impl RouteHook for SequencedHook {
async fn select(
&self,
_request: RouteRequest<'_>,
) -> std::result::Result<RouteDecision, RouteHookError> {
Ok(self
.decisions
.lock()
.expect("hook decisions mutex poisoned")
.remove(0))
}
}
struct RecordingHook {
decision: RouteDecision,
static_groups: Arc<Mutex<Vec<Option<UpstreamGroupId>>>>,
}
#[async_trait]
impl RouteHook for RecordingHook {
async fn select(
&self,
request: RouteRequest<'_>,
) -> std::result::Result<RouteDecision, RouteHookError> {
self.static_groups
.lock()
.expect("recorded static groups mutex poisoned")
.push(request.static_group().cloned());
Ok(self.decision.clone())
}
}
struct PendingHook {
entered: Arc<Notify>,
dropped: Arc<AtomicBool>,
}
struct DropSignal(Arc<AtomicBool>);
impl Drop for DropSignal {
fn drop(&mut self) {
self.0.store(true, Ordering::SeqCst);
}
}
#[async_trait]
impl RouteHook for PendingHook {
async fn select(
&self,
_request: RouteRequest<'_>,
) -> std::result::Result<RouteDecision, RouteHookError> {
let _drop_signal = DropSignal(self.dropped.clone());
self.entered.notify_waiters();
std::future::pending().await
}
}
fn a_answer(name: &str, ttl: u32) -> Message {
let mut msg = query_for(name);
msg.header.qr = true;
msg.answers.push(ResourceRecord {
name: n(name),
rtype: RecordType::A,
class: Class::In,
ttl,
rdata: RData::A(Ipv4Addr::new(203, 0, 113, 1)),
});
msg
}
fn nxdomain_answer(name: &str, soa_minimum: Option<u32>) -> Message {
let mut msg = query_for(name);
msg.header.qr = true;
msg.header.rcode = Rcode::NxDomain;
if let Some(minimum) = soa_minimum {
msg.authorities.push(ResourceRecord {
name: n("example.com"),
rtype: RecordType::Soa,
class: Class::In,
ttl: 3600,
rdata: RData::Soa {
mname: n("ns1.example.com"),
rname: n("hostmaster.example.com"),
serial: 1,
refresh: 3600,
retry: 600,
expire: 604_800,
minimum,
},
});
}
msg
}
fn nodata_answer(name: &str) -> Message {
query_for_response(name)
}
fn query_for_response(name: &str) -> Message {
let mut msg = query_for(name);
msg.header.qr = true;
msg
}
fn resolver_with_counting_backend(
policy: SplitDnsPolicy,
group: &str,
answer: Message,
clock: FakeClock,
) -> (Resolver, Arc<AtomicUsize>) {
let calls = Arc::new(AtomicUsize::new(0));
let backend = CountingBackend {
answer,
calls: calls.clone(),
};
let resolver = Resolver::builder(policy)
.clock(clock)
.backend(UpstreamGroupId::new(group), backend)
.build();
(resolver, calls)
}
#[tokio::test]
async fn cache_hit_does_not_call_backend_again() {
let policy = SplitDnsPolicy::builder()
.default_group(UpstreamGroupId::new("g"))
.build();
let (resolver, calls) = resolver_with_counting_backend(
policy,
"g",
a_answer("example.com", 300),
FakeClock::new(),
);
let first = resolver
.resolve(&query_for("example.com"))
.await
.expect("first resolve populates cache");
let second = resolver
.resolve(&query_for("example.com"))
.await
.expect("second resolve served from cache");
assert_eq!(first, second);
assert_eq!(calls.load(Ordering::SeqCst), 1, "backend called only once");
}
#[tokio::test]
async fn observability_reports_ordered_cache_miss_and_hit_without_affecting_resolution() {
let sink = Arc::new(RecordingSink::default());
let policy = SplitDnsPolicy::builder()
.default_group(UpstreamGroupId::new("g"))
.build();
let (base, calls) = resolver_with_counting_backend(
policy,
"g",
a_answer("example.com", 300),
FakeClock::new(),
);
let resolver = ResolverBuilder {
policy: base.policy,
backends: base.backends,
clock: base.clock,
fake_ip: base.fake_ip,
route_hook: base.route_hook,
observability_sink: Some(sink.clone()),
}
.build();
resolver.resolve(&query_for("example.com")).await.unwrap();
resolver.resolve(&query_for("example.com")).await.unwrap();
assert_eq!(calls.load(Ordering::SeqCst), 1);
let events = sink.0.lock().unwrap().clone();
assert!(matches!(
events[0],
ObserveEvent::QueryReceived {
correlation_id: 1,
..
}
));
assert!(matches!(
events[1],
ObserveEvent::StaticRoute {
correlation_id: 1,
..
}
));
assert!(matches!(
events[2],
ObserveEvent::CacheMiss {
correlation_id: 1,
..
}
));
assert!(matches!(
events[3],
ObserveEvent::UpstreamAttempt {
correlation_id: 1,
backend_index: 0,
..
}
));
assert!(matches!(
events[4],
ObserveEvent::UpstreamOutcome {
correlation_id: 1,
outcome: UpstreamObserveOutcome::Success,
..
}
));
assert!(matches!(
events[5],
ObserveEvent::Completed {
correlation_id: 1,
..
}
));
assert!(matches!(
events[6],
ObserveEvent::QueryReceived {
correlation_id: 2,
..
}
));
assert!(matches!(
events[7],
ObserveEvent::StaticRoute {
correlation_id: 2,
..
}
));
assert!(matches!(
events[8],
ObserveEvent::CacheHit {
correlation_id: 2,
..
}
));
assert!(matches!(
events[9],
ObserveEvent::Completed {
correlation_id: 2,
..
}
));
}
#[tokio::test]
async fn panicking_observability_sink_is_non_authoritative() {
let resolver = Resolver::builder(
SplitDnsPolicy::builder()
.default_group(UpstreamGroupId::new("g"))
.build(),
)
.backend(
UpstreamGroupId::new("g"),
fixed_backend(a_answer("example.com", 300)),
)
.observability_sink(Arc::new(PanickingSink))
.build();
assert!(resolver.resolve(&query_for("example.com")).await.is_ok());
}
#[tokio::test]
async fn observability_starts_empty_queries_before_no_route_failure() {
let sink = Arc::new(RecordingSink::default());
let resolver = Resolver::builder(SplitDnsPolicy::builder().build())
.observability_sink(sink.clone())
.build();
let mut query = query_for("example.com");
query.questions.clear();
assert_eq!(resolver.resolve(&query).await, Err(Error::NoRoute));
assert!(matches!(
sink.0.lock().unwrap().as_slice(),
[
ObserveEvent::QueryReceived {
name: None,
rtype: None,
class: None,
..
},
ObserveEvent::Failed {
failure: ObserveFailure::NoRoute,
..
},
]
));
}
#[tokio::test]
async fn observability_marks_fake_ip_terminal_before_cache_or_upstream() {
let sink = Arc::new(RecordingSink::default());
let resolver = Resolver::builder(SplitDnsPolicy::builder().build())
.fake_ip(
fake_ip_pool(PoolClock::new()),
fake_ip_policy("example.com"),
)
.observability_sink(sink.clone())
.build();
resolver.resolve(&query_for("example.com")).await.unwrap();
assert!(matches!(
sink.0.lock().unwrap().as_slice(),
[
ObserveEvent::QueryReceived { .. },
ObserveEvent::FakeIpTerminal { .. },
ObserveEvent::Completed { .. },
]
));
}
#[tokio::test]
async fn observability_records_hook_and_timeout_failures_in_order() {
let sink = Arc::new(RecordingSink::default());
let resolver = Resolver::builder(
SplitDnsPolicy::builder()
.default_group(UpstreamGroupId::new("g"))
.build(),
)
.route_hook(FixedHook {
decision: Err(RouteHookError::new("denied")),
calls: Arc::new(AtomicUsize::new(0)),
})
.observability_sink(sink.clone())
.build();
assert!(matches!(
resolver.resolve(&query_for("example.com")).await,
Err(Error::Hook(_))
));
assert!(matches!(
sink.0.lock().unwrap().as_slice(),
[
ObserveEvent::QueryReceived { .. },
ObserveEvent::StaticRoute { .. },
ObserveEvent::HookDecision {
decision: HookObserveDecision::Failed,
..
},
ObserveEvent::Failed {
failure: ObserveFailure::Hook,
..
},
]
));
sink.0.lock().unwrap().clear();
let resolver = Resolver::builder(
SplitDnsPolicy::builder()
.default_group(UpstreamGroupId::new("g"))
.build(),
)
.backend(UpstreamGroupId::new("g"), FailingBackend(Error::Timeout))
.observability_sink(sink.clone())
.build();
assert_eq!(
resolver.resolve(&query_for("example.com")).await,
Err(Error::Timeout)
);
assert!(matches!(
sink.0.lock().unwrap().as_slice(),
[
ObserveEvent::QueryReceived { .. },
ObserveEvent::StaticRoute { .. },
ObserveEvent::CacheMiss { .. },
ObserveEvent::UpstreamAttempt { .. },
ObserveEvent::UpstreamOutcome {
outcome: UpstreamObserveOutcome::RetryableFailure,
..
},
ObserveEvent::Failed {
failure: ObserveFailure::Timeout,
..
},
]
));
}
#[tokio::test]
async fn cache_hit_preserves_the_current_query_identity_and_questions() {
let policy = SplitDnsPolicy::builder()
.default_group(UpstreamGroupId::new("g"))
.build();
let (resolver, calls) = resolver_with_counting_backend(
policy,
"g",
a_answer("example.com", 300),
FakeClock::new(),
);
let first = query_for_type("example.com", RecordType::A, Class::In, 91);
let mut second = query_for_type("example.com", RecordType::A, Class::In, 92);
second.questions.push(Question {
name: n("extra.example.com"),
qtype: RecordType::Aaaa,
qclass: Class::In,
});
resolver
.resolve(&first)
.await
.expect("first resolve populates cache");
let cached = resolver
.resolve(&second)
.await
.expect("second resolve is served from cache");
assert_eq!(cached.header.id, 92);
assert_eq!(cached.questions, second.questions);
assert_eq!(
cached.answers[0].rdata,
RData::A(Ipv4Addr::new(203, 0, 113, 1))
);
assert_eq!(
calls.load(Ordering::SeqCst),
1,
"second query is a cache hit"
);
}
#[tokio::test]
async fn cache_identity_separates_generated_question_type_and_class_pairs() {
let policy = SplitDnsPolicy::builder()
.default_group(UpstreamGroupId::new("g"))
.build();
let (resolver, calls) = resolver_with_counting_backend(
policy,
"g",
a_answer("example.com", 300),
FakeClock::new(),
);
let cases = [
(RecordType::A, Class::In, 1),
(RecordType::Aaaa, Class::In, 2),
(RecordType::A, Class::Ch, 3),
(RecordType::Other(65280), Class::Other(65280), 4),
];
for (rtype, class, id) in cases {
resolver
.resolve(&query_for_type("example.com", rtype, class, id))
.await
.expect("each distinct cache identity resolves");
}
assert_eq!(calls.load(Ordering::SeqCst), cases.len());
for (rtype, class, id) in cases {
let cached = resolver
.resolve(&query_for_type("example.com", rtype, class, id + 10))
.await
.expect("same type/class pair is cached");
assert_eq!(cached.header.id, id + 10);
assert_eq!(cached.questions[0].qtype, rtype);
assert_eq!(cached.questions[0].qclass, class);
}
assert_eq!(calls.load(Ordering::SeqCst), cases.len());
}
#[tokio::test]
async fn cache_entry_still_hit_just_before_ttl_elapses() {
let policy = SplitDnsPolicy::builder()
.default_group(UpstreamGroupId::new("g"))
.build();
let clock = FakeClock::new();
let (resolver, calls) = resolver_with_counting_backend(
policy,
"g",
a_answer("example.com", 300),
clock.clone(),
);
resolver
.resolve(&query_for("example.com"))
.await
.expect("first resolve populates cache");
assert_eq!(calls.load(Ordering::SeqCst), 1);
clock.advance(Duration::from_secs(299));
resolver
.resolve(&query_for("example.com"))
.await
.expect("still cached before ttl elapses");
assert_eq!(calls.load(Ordering::SeqCst), 1, "cache hit before expiry");
}
#[tokio::test]
async fn negative_answer_is_cached_with_soa_minimum_ttl() {
let policy = SplitDnsPolicy::builder()
.default_group(UpstreamGroupId::new("g"))
.build();
let (resolver, calls) = resolver_with_counting_backend(
policy,
"g",
nxdomain_answer("missing.example.com", Some(300)),
FakeClock::new(),
);
let first = resolver
.resolve(&query_for("missing.example.com"))
.await
.expect("nxdomain is Ok(Message), not Err");
assert_eq!(first.header.rcode, Rcode::NxDomain);
resolver
.resolve(&query_for("missing.example.com"))
.await
.expect("served from negative cache");
assert_eq!(calls.load(Ordering::SeqCst), 1, "negative answer cached");
}
#[tokio::test]
async fn negative_answer_without_soa_uses_fixed_floor_ttl() {
let policy = SplitDnsPolicy::builder()
.default_group(UpstreamGroupId::new("g"))
.build();
let (resolver, calls) = resolver_with_counting_backend(
policy,
"g",
nxdomain_answer("missing.example.com", None),
FakeClock::new(),
);
resolver
.resolve(&query_for("missing.example.com"))
.await
.expect("nxdomain without soa still Ok");
resolver
.resolve(&query_for("missing.example.com"))
.await
.expect("served from cache using the fixed floor ttl");
assert_eq!(
calls.load(Ordering::SeqCst),
1,
"negative answer cached via floor"
);
}
#[tokio::test]
async fn nodata_answer_is_cached_as_negative() {
let policy = SplitDnsPolicy::builder()
.default_group(UpstreamGroupId::new("g"))
.build();
let (resolver, calls) = resolver_with_counting_backend(
policy,
"g",
nodata_answer("empty.example.com"),
FakeClock::new(),
);
resolver
.resolve(&query_for("empty.example.com"))
.await
.expect("nodata is Ok(Message)");
resolver
.resolve(&query_for("empty.example.com"))
.await
.expect("served from cache");
assert_eq!(calls.load(Ordering::SeqCst), 1, "nodata answer cached");
}
#[tokio::test]
async fn expired_cache_entry_triggers_a_fresh_backend_call() {
let policy = SplitDnsPolicy::builder()
.default_group(UpstreamGroupId::new("g"))
.build();
let clock = FakeClock::new();
let (resolver, calls) =
resolver_with_counting_backend(policy, "g", a_answer("example.com", 10), clock.clone());
resolver
.resolve(&query_for("example.com"))
.await
.expect("first resolve populates cache");
assert_eq!(calls.load(Ordering::SeqCst), 1);
clock.advance(Duration::from_secs(11));
resolver
.resolve(&query_for("example.com"))
.await
.expect("expired entry re-queries the backend");
assert_eq!(
calls.load(Ordering::SeqCst),
2,
"ttl-expired entry is not served from cache"
);
}
#[tokio::test]
async fn hook_use_overrides_the_static_group() {
let hook_calls = Arc::new(AtomicUsize::new(0));
let static_calls = Arc::new(AtomicUsize::new(0));
let selected_calls = Arc::new(AtomicUsize::new(0));
let resolver = Resolver::builder(
SplitDnsPolicy::builder()
.default_group(UpstreamGroupId::new("static"))
.build(),
)
.backend(
UpstreamGroupId::new("static"),
CountingBackend {
answer: answer_tagged(1),
calls: static_calls.clone(),
},
)
.backend(
UpstreamGroupId::new("selected"),
CountingBackend {
answer: answer_tagged(2),
calls: selected_calls.clone(),
},
)
.route_hook(FixedHook {
decision: Ok(RouteDecision::Use(UpstreamGroupId::new("selected"))),
calls: hook_calls.clone(),
})
.build();
let answer = resolver.resolve(&query_for("example.com")).await.unwrap();
assert_eq!(answer.header.id, 2);
assert_eq!(hook_calls.load(Ordering::SeqCst), 1);
assert_eq!(static_calls.load(Ordering::SeqCst), 0);
assert_eq!(selected_calls.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn hook_abstain_uses_the_static_group() {
let backend_calls = Arc::new(AtomicUsize::new(0));
let resolver = Resolver::builder(
SplitDnsPolicy::builder()
.default_group(UpstreamGroupId::new("static"))
.build(),
)
.backend(
UpstreamGroupId::new("static"),
CountingBackend {
answer: answer_tagged(3),
calls: backend_calls.clone(),
},
)
.route_hook(FixedHook {
decision: Ok(RouteDecision::Abstain),
calls: Arc::new(AtomicUsize::new(0)),
})
.build();
assert_eq!(
resolver
.resolve(&query_for("example.com"))
.await
.unwrap()
.header
.id,
3
);
assert_eq!(backend_calls.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn hook_observes_static_candidate_and_can_supply_a_route_without_one() {
let static_groups = Arc::new(Mutex::new(Vec::new()));
let static_resolver = Resolver::builder(
SplitDnsPolicy::builder()
.default_group(UpstreamGroupId::new("static"))
.build(),
)
.backend(
UpstreamGroupId::new("static"),
fixed_backend(answer_tagged(30)),
)
.route_hook(RecordingHook {
decision: RouteDecision::Abstain,
static_groups: static_groups.clone(),
})
.build();
assert_eq!(
static_resolver
.resolve(&query_for("static.example"))
.await
.unwrap()
.header
.id,
30
);
let dynamic_resolver = Resolver::builder(SplitDnsPolicy::builder().build())
.backend(
UpstreamGroupId::new("dynamic"),
fixed_backend(answer_tagged(31)),
)
.route_hook(RecordingHook {
decision: RouteDecision::Use(UpstreamGroupId::new("dynamic")),
static_groups: static_groups.clone(),
})
.build();
assert_eq!(
dynamic_resolver
.resolve(&query_for("dynamic.example"))
.await
.unwrap()
.header
.id,
31
);
assert_eq!(
*static_groups.lock().unwrap(),
vec![Some(UpstreamGroupId::new("static")), None]
);
}
#[tokio::test]
async fn hook_abstain_without_static_route_returns_no_route() {
let backend_calls = Arc::new(AtomicUsize::new(0));
let resolver = Resolver::builder(SplitDnsPolicy::builder().build())
.backend(
UpstreamGroupId::new("unused"),
CountingBackend {
answer: answer_tagged(4),
calls: backend_calls.clone(),
},
)
.route_hook(FixedHook {
decision: Ok(RouteDecision::Abstain),
calls: Arc::new(AtomicUsize::new(0)),
})
.build();
assert_eq!(
resolver.resolve(&query_for("example.com")).await,
Err(Error::NoRoute)
);
assert_eq!(backend_calls.load(Ordering::SeqCst), 0);
}
#[tokio::test]
async fn hook_selected_unknown_or_empty_group_returns_no_route_without_fallback() {
for group in ["unknown", "empty"] {
let static_calls = Arc::new(AtomicUsize::new(0));
let builder = Resolver::builder(
SplitDnsPolicy::builder()
.default_group(UpstreamGroupId::new("static"))
.build(),
)
.backend(
UpstreamGroupId::new("static"),
CountingBackend {
answer: answer_tagged(5),
calls: static_calls.clone(),
},
);
let mut resolver = builder
.route_hook(FixedHook {
decision: Ok(RouteDecision::Use(UpstreamGroupId::new(group))),
calls: Arc::new(AtomicUsize::new(0)),
})
.build();
if group == "empty" {
resolver
.backends
.insert(UpstreamGroupId::new("empty"), Vec::new());
}
assert_eq!(
resolver.resolve(&query_for("example.com")).await,
Err(Error::NoRoute)
);
assert_eq!(
static_calls.load(Ordering::SeqCst),
0,
"static backend must not receive a hook-selected {group} route"
);
}
}
#[tokio::test]
async fn hook_error_is_not_cached_retried_or_fallen_back() {
let hook_calls = Arc::new(AtomicUsize::new(0));
let backend_calls = Arc::new(AtomicUsize::new(0));
let resolver = Resolver::builder(
SplitDnsPolicy::builder()
.default_group(UpstreamGroupId::new("static"))
.build(),
)
.backend(
UpstreamGroupId::new("static"),
CountingBackend {
answer: answer_tagged(6),
calls: backend_calls.clone(),
},
)
.route_hook(FixedHook {
decision: Err(RouteHookError::new("policy unavailable")),
calls: hook_calls.clone(),
})
.build();
for _ in 0..2 {
assert_eq!(
resolver.resolve(&query_for("example.com")).await,
Err(Error::Hook("policy unavailable".to_string()))
);
}
assert_eq!(hook_calls.load(Ordering::SeqCst), 2);
assert_eq!(backend_calls.load(Ordering::SeqCst), 0);
assert!(resolver.cache.lock().unwrap().is_empty());
}
#[tokio::test]
async fn cache_is_scoped_to_the_effective_hook_selected_group() {
let first_calls = Arc::new(AtomicUsize::new(0));
let second_calls = Arc::new(AtomicUsize::new(0));
let resolver = Resolver::builder(SplitDnsPolicy::builder().build())
.backend(
UpstreamGroupId::new("first"),
CountingBackend {
answer: a_answer("example.com", 300),
calls: first_calls.clone(),
},
)
.backend(
UpstreamGroupId::new("second"),
CountingBackend {
answer: answer_tagged(8),
calls: second_calls.clone(),
},
)
.route_hook(SequencedHook {
decisions: Mutex::new(vec![
RouteDecision::Use(UpstreamGroupId::new("first")),
RouteDecision::Use(UpstreamGroupId::new("second")),
RouteDecision::Use(UpstreamGroupId::new("first")),
]),
})
.build();
let first = resolver
.resolve(&query_for_type("example.com", RecordType::A, Class::In, 41))
.await
.unwrap();
let second = resolver
.resolve(&query_for_type("example.com", RecordType::A, Class::In, 42))
.await
.unwrap();
let cached_first = resolver
.resolve(&query_for_type("example.com", RecordType::A, Class::In, 43))
.await
.unwrap();
assert_eq!(first.answers[0].ttl, 300);
assert_eq!(
second.header.id, 8,
"second route cannot reuse first route cache"
);
assert_eq!(cached_first.header.id, 43);
assert_eq!(cached_first.questions, query_for("example.com").questions);
assert_eq!(
cached_first.answers, first.answers,
"first route has its own cache hit"
);
assert_eq!(first_calls.load(Ordering::SeqCst), 1);
assert_eq!(second_calls.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn dropping_resolve_drops_the_hook_future_without_holding_cache_lock() {
let entered = Arc::new(Notify::new());
let dropped = Arc::new(AtomicBool::new(false));
let resolver = Arc::new(
Resolver::builder(SplitDnsPolicy::builder().build())
.route_hook(PendingHook {
entered: entered.clone(),
dropped: dropped.clone(),
})
.build(),
);
let entered_wait = entered.notified();
let task_resolver = resolver.clone();
let task =
tokio::spawn(async move { task_resolver.resolve(&query_for("example.com")).await });
entered_wait.await;
assert!(
resolver.cache.try_lock().is_ok(),
"the resolver cache mutex is not held across hook await"
);
task.abort();
assert!(task.await.unwrap_err().is_cancelled());
assert!(
dropped.load(Ordering::SeqCst),
"hook future was dropped on cancellation"
);
}
fn query_for_type(name: &str, qtype: RecordType, qclass: Class, id: u16) -> Message {
let mut query = query_for(name);
query.header.id = id;
query.questions[0].qtype = qtype;
query.questions[0].qclass = qclass;
query
}
fn fake_ip_policy(name: &str) -> FakeIpPolicy {
FakeIpPolicy::builder()
.rule(DomainPattern::suffix(n(name)))
.build()
}
fn fake_ip_pool(clock: PoolClock) -> Arc<FakeIpPool> {
Arc::new(
FakeIpPool::builder()
.ipv4_range(Ipv4Addr::new(198, 18, 0, 1), Ipv4Addr::new(198, 18, 0, 2))
.ttl(Duration::from_secs(30))
.clock(clock)
.build()
.unwrap(),
)
}
fn fake_ip_pool_ipv6(clock: PoolClock) -> Arc<FakeIpPool> {
Arc::new(
FakeIpPool::builder()
.ipv6_range(
"2001:db8::1".parse().unwrap(),
"2001:db8::2".parse().unwrap(),
)
.ttl(Duration::from_secs(30))
.clock(clock)
.build()
.unwrap(),
)
}
#[tokio::test]
async fn fake_ip_a_answer_is_local_and_bypasses_upstream_and_cache() {
let calls = Arc::new(AtomicUsize::new(0));
let hook_calls = Arc::new(AtomicUsize::new(0));
let backend = CountingBackend {
answer: a_answer("example.test", 300),
calls: calls.clone(),
};
let pool = fake_ip_pool(PoolClock::new());
let resolver = Resolver::builder(
SplitDnsPolicy::builder()
.default_group(UpstreamGroupId::new("g"))
.build(),
)
.backend(UpstreamGroupId::new("g"), backend)
.fake_ip(pool, fake_ip_policy("example.test"))
.route_hook(FixedHook {
decision: Ok(RouteDecision::Use(UpstreamGroupId::new("g"))),
calls: hook_calls.clone(),
})
.build();
let first = resolver
.resolve(&query_for_type(
"www.example.test",
RecordType::A,
Class::In,
41,
))
.await
.unwrap();
let second = resolver
.resolve(&query_for_type(
"www.example.test",
RecordType::A,
Class::In,
42,
))
.await
.unwrap();
assert_eq!(calls.load(Ordering::SeqCst), 0);
assert_eq!(
hook_calls.load(Ordering::SeqCst),
0,
"Fake IP is terminal before hooks"
);
assert_eq!(first.header.id, 41);
assert_eq!(second.header.id, 42, "synthetic answers are not cached");
assert!(first.header.qr);
assert_eq!(first.questions, query_for("www.example.test").questions);
assert_eq!(first.answers[0].ttl, 30);
assert_eq!(first.answers[0].rdata, second.answers[0].rdata);
}
#[tokio::test]
async fn fake_ip_disabled_family_returns_local_nodata() {
let calls = Arc::new(AtomicUsize::new(0));
let resolver = Resolver::builder(
SplitDnsPolicy::builder()
.default_group(UpstreamGroupId::new("g"))
.build(),
)
.backend(
UpstreamGroupId::new("g"),
CountingBackend {
answer: a_answer("example.test", 300),
calls: calls.clone(),
},
)
.fake_ip(
fake_ip_pool(PoolClock::new()),
fake_ip_policy("example.test"),
)
.build();
let answer = resolver
.resolve(&query_for_type(
"www.example.test",
RecordType::Aaaa,
Class::In,
9,
))
.await
.unwrap();
assert_eq!(answer.header.rcode, Rcode::NoError);
assert!(answer.answers.is_empty());
assert_eq!(calls.load(Ordering::SeqCst), 0);
}
#[tokio::test]
async fn fake_ip_ptr_is_local_and_expires_with_its_mapping() {
let pool_clock = PoolClock::new();
let pool = fake_ip_pool(pool_clock.clone());
let resolver = Resolver::builder(SplitDnsPolicy::builder().build())
.fake_ip(pool.clone(), fake_ip_policy("example.test"))
.build();
let address = pool.allocate_ipv4(n("www.example.test")).unwrap();
let reverse = format!(
"{}.{}.{}.{}.in-addr.arpa",
address.octets()[3],
address.octets()[2],
address.octets()[1],
address.octets()[0]
);
let found = resolver
.resolve(&query_for_type(&reverse, RecordType::Ptr, Class::In, 11))
.await
.unwrap();
assert_eq!(found.header.rcode, Rcode::NoError);
assert_eq!(found.answers[0].rdata, RData::Ptr(n("www.example.test")));
assert_eq!(found.answers[0].ttl, 30);
pool_clock.advance(Duration::from_secs(30));
let expired = resolver
.resolve(&query_for_type(&reverse, RecordType::Ptr, Class::In, 12))
.await
.unwrap();
assert_eq!(expired.header.rcode, Rcode::NxDomain);
assert!(expired.answers.is_empty());
}
#[tokio::test]
async fn fake_ip_answer_ttl_never_outlives_existing_mapping() {
let pool_clock = PoolClock::new();
let pool = fake_ip_pool(pool_clock.clone());
let resolver = Resolver::builder(SplitDnsPolicy::builder().build())
.fake_ip(pool.clone(), fake_ip_policy("example.test"))
.build();
pool.allocate_ipv4(n("www.example.test")).unwrap();
pool_clock.advance(Duration::from_secs(29));
let answer = resolver
.resolve(&query_for_type(
"www.example.test",
RecordType::A,
Class::In,
20,
))
.await
.unwrap();
assert_eq!(answer.answers[0].ttl, 1);
}
#[tokio::test]
async fn fake_ip_ptr_ttl_never_outlives_existing_mapping() {
let pool_clock = PoolClock::new();
let pool = fake_ip_pool(pool_clock.clone());
let resolver = Resolver::builder(SplitDnsPolicy::builder().build())
.fake_ip(pool.clone(), fake_ip_policy("example.test"))
.build();
let address = pool.allocate_ipv4(n("www.example.test")).unwrap();
let reverse = format!(
"{}.{}.{}.{}.in-addr.arpa",
address.octets()[3],
address.octets()[2],
address.octets()[1],
address.octets()[0]
);
pool_clock.advance(Duration::from_secs(29));
let answer = resolver
.resolve(&query_for_type(&reverse, RecordType::Ptr, Class::In, 21))
.await
.unwrap();
assert_eq!(answer.answers[0].ttl, 1);
}
#[tokio::test]
async fn fake_ip_ipv6_ptr_is_local() {
let pool = fake_ip_pool_ipv6(PoolClock::new());
let resolver = Resolver::builder(SplitDnsPolicy::builder().build())
.fake_ip(pool.clone(), fake_ip_policy("example.test"))
.build();
let address = pool.allocate_ipv6(n("www.example.test")).unwrap();
let reverse = address
.octets()
.iter()
.rev()
.flat_map(|byte| [format!("{:x}", byte & 0x0f), format!("{:x}", byte >> 4)])
.collect::<Vec<_>>()
.join(".");
let answer = resolver
.resolve(&query_for_type(
&format!("{reverse}.ip6.arpa"),
RecordType::Ptr,
Class::In,
22,
))
.await
.unwrap();
assert_eq!(answer.header.rcode, Rcode::NoError);
assert_eq!(answer.answers[0].rdata, RData::Ptr(n("www.example.test")));
}
#[tokio::test]
async fn normal_queries_and_outside_reverse_ranges_still_use_upstream() {
let calls = Arc::new(AtomicUsize::new(0));
let pool = Arc::new(
FakeIpPool::builder()
.ipv4_range(Ipv4Addr::new(198, 18, 0, 1), Ipv4Addr::new(198, 18, 0, 2))
.ttl(Duration::from_secs(u64::MAX))
.clock(PoolClock::new())
.build()
.unwrap(),
);
let resolver = Resolver::builder(
SplitDnsPolicy::builder()
.default_group(UpstreamGroupId::new("g"))
.build(),
)
.backend(
UpstreamGroupId::new("g"),
CountingBackend {
answer: answer_tagged(77),
calls: calls.clone(),
},
)
.fake_ip(pool, fake_ip_policy("selected.test"))
.build();
for query in [
query_for_type("miss.test", RecordType::A, Class::In, 1),
query_for_type("selected.test", RecordType::A, Class::Ch, 2),
query_for_type("selected.test", RecordType::Txt, Class::In, 3),
query_for_type("1.0.0.203.in-addr.arpa", RecordType::Ptr, Class::In, 4),
] {
let answer = resolver.resolve(&query).await.unwrap();
assert_eq!(answer.header.id, 77);
}
assert_eq!(calls.load(Ordering::SeqCst), 4);
}
#[tokio::test]
async fn unrepresentable_fake_ip_ttl_fails_before_allocation() {
let pool = Arc::new(
FakeIpPool::builder()
.ipv4_range(Ipv4Addr::new(198, 18, 0, 1), Ipv4Addr::new(198, 18, 0, 2))
.ttl(Duration::from_secs(u64::from(u32::MAX) + 1))
.clock(PoolClock::new())
.build()
.unwrap(),
);
let resolver = Resolver::builder(SplitDnsPolicy::builder().build())
.fake_ip(pool.clone(), fake_ip_policy("example.test"))
.build();
assert_eq!(
resolver
.resolve(&query_for_type(
"www.example.test",
RecordType::A,
Class::In,
23
))
.await,
Err(Error::FakeIpTtlOutOfRange)
);
assert!(pool.snapshot().mappings.is_empty());
}
#[tokio::test]
async fn disabled_fake_ip_families_return_nodata_even_with_unrepresentable_ttl() {
let calls = Arc::new(AtomicUsize::new(0));
let pool = Arc::new(
FakeIpPool::builder()
.ipv4_range(Ipv4Addr::new(198, 18, 0, 1), Ipv4Addr::new(198, 18, 0, 2))
.ttl(Duration::from_secs(u64::from(u32::MAX) + 1))
.clock(PoolClock::new())
.build()
.unwrap(),
);
let resolver = Resolver::builder(
SplitDnsPolicy::builder()
.default_group(UpstreamGroupId::new("g"))
.build(),
)
.backend(
UpstreamGroupId::new("g"),
CountingBackend {
answer: answer_tagged(78),
calls: calls.clone(),
},
)
.fake_ip(pool.clone(), fake_ip_policy("example.test"))
.build();
let aaaa = resolver
.resolve(&query_for_type(
"www.example.test",
RecordType::Aaaa,
Class::In,
24,
))
.await
.unwrap();
assert_eq!(aaaa.header.rcode, Rcode::NoError);
assert!(aaaa.answers.is_empty());
assert!(pool.snapshot().mappings.is_empty());
assert_eq!(calls.load(Ordering::SeqCst), 0);
let ipv6_only_pool = Arc::new(
FakeIpPool::builder()
.ipv6_range(
"2001:db8::1".parse().unwrap(),
"2001:db8::2".parse().unwrap(),
)
.ttl(Duration::from_secs(u64::from(u32::MAX) + 1))
.clock(PoolClock::new())
.build()
.unwrap(),
);
let ipv6_only_resolver = Resolver::builder(
SplitDnsPolicy::builder()
.default_group(UpstreamGroupId::new("g"))
.build(),
)
.backend(
UpstreamGroupId::new("g"),
CountingBackend {
answer: answer_tagged(79),
calls: calls.clone(),
},
)
.fake_ip(ipv6_only_pool.clone(), fake_ip_policy("example.test"))
.build();
let a = ipv6_only_resolver
.resolve(&query_for_type(
"www.example.test",
RecordType::A,
Class::In,
25,
))
.await
.unwrap();
assert_eq!(a.header.rcode, Rcode::NoError);
assert!(a.answers.is_empty());
assert!(ipv6_only_pool.snapshot().mappings.is_empty());
assert_eq!(calls.load(Ordering::SeqCst), 0);
}
#[test]
fn parses_canonical_ipv4_and_ipv6_reverse_names() {
assert_eq!(
parse_reverse_name(&n("4.3.2.1.in-addr.arpa")),
Some(std::net::IpAddr::V4(Ipv4Addr::new(1, 2, 3, 4)))
);
assert_eq!(
parse_reverse_name(&n("4.3.2.01.in-addr.arpa")),
None,
"non-canonical decimal labels are routed normally"
);
let reverse = format!("1.{}ip6.arpa", "0.".repeat(31));
assert_eq!(
parse_reverse_name(&n(&reverse)),
Some(std::net::IpAddr::V6("::1".parse().unwrap()))
);
}
}