Skip to main content

dns_lattice/
engine.rs

1//! Query orchestration for decoded DNS messages.
2//!
3//! [`Resolver`] accepts a decoded [`Message`], selects an upstream group by
4//! static [`SplitDnsPolicy`] routing, reads and writes its in-memory
5//! TTL/negative cache, and invokes registered [`crate::upstream::UpstreamBackend`]
6//! values in registration order with retryable-error failover.
7//!
8//! It does **not** own inbound server lifecycle, socket binding, wire
9//! framing, TLS/HTTP/QUIC protocol handling, operating-system DNS
10//! configuration, or packet forwarding. Those responsibilities belong
11//! respectively to [`crate::server`], [`crate::upstream`], and composing
12//! applications. An optional [`crate::hooks::RouteHook`] selects an existing
13//! upstream group; it does not own resolution or side effects. When
14//! explicitly configured with a
15//! [`crate::fakeip::FakeIpPool`] and [`crate::fakeip::FakeIpPolicy`], it does
16//! orchestrate their local DNS answer synthesis; allocation and mapping
17//! storage remain owned by the pool.
18
19use std::collections::HashMap;
20use std::net::{Ipv4Addr, Ipv6Addr};
21use std::sync::atomic::{AtomicU64, Ordering};
22use std::sync::{Arc, Mutex};
23use std::time::{Duration, Instant};
24
25use dns_lattice_core::{Error, Result};
26use dns_lattice_model::{
27    Class, Message, Name, RData, Rcode, RecordType, ResourceRecord, SplitDnsPolicy, UpstreamGroupId,
28};
29
30use crate::fakeip::{FakeIpPolicy, FakeIpPool};
31use crate::hooks::{RouteDecision, RouteHook, RouteRequest};
32use crate::observability::{
33    HookObserveDecision, ObservabilitySink, ObserveEvent, ObserveFailure, UpstreamObserveOutcome,
34};
35use crate::upstream::UpstreamBackend;
36
37/// Fixed negative-cache TTL floor (RFC 2308 ยง5) used when a negative
38/// response carries no SOA record in its authority section to derive a
39/// `minimum` from. It is not user-configurable.
40const NEGATIVE_CACHE_FLOOR: Duration = Duration::from_secs(60);
41
42/// A source of the current time, abstracted so tests can advance it
43/// deterministically instead of relying on real `sleep`.
44///
45/// Crate-private: no external caller needs to inject a clock in this stage;
46/// [`Resolver::builder`] always defaults to [`SystemClock`].
47pub(crate) trait Clock {
48    /// Returns the current instant.
49    fn now(&self) -> Instant;
50}
51
52/// Production [`Clock`] delegating to [`Instant::now`].
53pub(crate) struct SystemClock;
54
55impl Clock for SystemClock {
56    fn now(&self) -> Instant {
57        Instant::now()
58    }
59}
60
61/// A manually-advanced [`Clock`] for deterministic tests.
62///
63/// Interior-mutable and cheaply cloneable (shares the same underlying cell
64/// via `Arc<Mutex<_>>`, kept `Send + Sync` so it satisfies
65/// [`ResolverBuilder::clock`]'s bound) so a test can keep a handle to
66/// advance the clock after handing an owned copy to the resolver.
67#[cfg(test)]
68#[derive(Clone)]
69pub(crate) struct FakeClock(std::sync::Arc<std::sync::Mutex<Instant>>);
70
71#[cfg(test)]
72impl FakeClock {
73    /// Starts the clock at the current real instant (only used as an
74    /// arbitrary, non-real-time-dependent base point).
75    pub(crate) fn new() -> Self {
76        FakeClock(std::sync::Arc::new(std::sync::Mutex::new(Instant::now())))
77    }
78
79    /// Advances the clock by `duration`.
80    pub(crate) fn advance(&self, duration: Duration) {
81        let mut guard = self.0.lock().expect("fake clock mutex poisoned");
82        *guard += duration;
83    }
84}
85
86#[cfg(test)]
87impl Clock for FakeClock {
88    fn now(&self) -> Instant {
89        *self.0.lock().expect("fake clock mutex poisoned")
90    }
91}
92
93/// Cache key: the fields that identify a question's matching intent,
94/// equivalent to a [`dns_lattice_model::Question`]'s name/type/class but
95/// independent of that struct's exact field set.
96#[derive(Clone, PartialEq, Eq, Hash)]
97struct CacheKey {
98    name: Name,
99    rtype: RecordType,
100    class: Class,
101    group: UpstreamGroupId,
102}
103
104/// A cached answer plus its absolute expiry instant, computed at insert
105/// time.
106struct CacheEntry {
107    answer: Message,
108    expires_at: Instant,
109}
110
111/// An in-process DNS query orchestrator.
112///
113/// Construct it from a split-DNS policy and one or more upstream backends
114/// per group, then resolve decoded queries against it. It owns policy
115/// selection, caching, and upstream failover, but not server lifecycle or
116/// transport protocol implementation; see the [module documentation].
117///
118/// [module documentation]: self
119///
120/// # Lifecycle
121///
122/// Construct via [`Resolver::builder`], call [`Resolver::resolve`] as many
123/// times as needed, then drop. This stage holds no background threads, so
124/// there is no explicit `shutdown` method โ€” Rust's ordinary drop semantics
125/// fully release any resources the resolver owns (including any sockets a
126/// registered [`crate::upstream::UdpBackend`]/[`crate::upstream::TcpBackend`]
127/// opens per call).
128pub struct Resolver {
129    policy: SplitDnsPolicy,
130    backends: HashMap<UpstreamGroupId, Vec<Box<dyn UpstreamBackend>>>,
131    clock: Box<dyn Clock + Send + Sync>,
132    cache: Mutex<HashMap<CacheKey, CacheEntry>>,
133    fake_ip: Option<FakeIpResolverConfig>,
134    route_hook: Option<Box<dyn RouteHook>>,
135    observability_sink: Option<Arc<dyn ObservabilitySink>>,
136    next_correlation_id: AtomicU64,
137}
138
139/// Explicit Fake IP answer synthesis owned by a [`Resolver`].
140///
141/// The pool remains caller-owned through [`Arc`], allowing the composing
142/// application to perform lookups and retain mappings independently of this
143/// resolver. This configuration is deliberately opt-in.
144struct FakeIpResolverConfig {
145    pool: Arc<FakeIpPool>,
146    policy: FakeIpPolicy,
147}
148
149impl Resolver {
150    /// Starts building a resolver from a split-DNS policy.
151    pub fn builder(policy: SplitDnsPolicy) -> ResolverBuilder {
152        ResolverBuilder {
153            policy,
154            backends: HashMap::new(),
155            clock: Box::new(SystemClock),
156            fake_ip: None,
157            route_hook: None,
158            observability_sink: None,
159        }
160    }
161
162    /// Resolves one query.
163    ///
164    /// Extracts the queried name from `query`'s first question. Locally
165    /// handled Fake IP questions return before the hook, cache, and upstream
166    /// stages. For every other question, static [`SplitDnsPolicy`] routing
167    /// supplies a tentative group and an optional [`crate::hooks::RouteHook`]
168    /// can authoritatively replace it. The selected registered group scopes
169    /// the in-memory answer cache; on a miss its backends are tried in
170    /// registration order: the first
171    /// backend to return `Ok` wins and its answer is
172    /// cached (per the existing TTL rules) and returned immediately. A
173    /// backend failing with [`Error::Timeout`], [`Error::Transport`], or
174    /// [`Error::Tls`] is treated as retryable โ€” resolution moves on to the
175    /// next backend in the group rather than failing the whole call. Once
176    /// every backend in the group has been tried and failed, the *last*
177    /// attempted backend's error is propagated as-is; this exhausted-group
178    /// failure is never cached. A group with exactly one backend behaves
179    /// exactly as before: success or that one backend's own error.
180    ///
181    /// # Errors
182    ///
183    /// Returns [`Error::NoRoute`] when no split-DNS rule matches the queried
184    /// name and no default upstream group is configured, when no question is
185    /// present in `query`, or when the selected group has no backend
186    /// registered. A hook-selected unknown or empty group never falls back to
187    /// static policy. A hook failure returns [`Error::Hook`] and is neither
188    /// retried nor cached. Returns the last attempted backend's error,
189    /// propagated as-is (not cached), once every backend in the matched
190    /// group has failed.
191    ///
192    /// # Runtime requirement
193    ///
194    /// Must be called from inside a `tokio` runtime context if the selected
195    /// backend performs real socket I/O (e.g. [`crate::upstream::UdpBackend`]/
196    /// [`crate::upstream::TcpBackend`]) โ€” see `crate::upstream`'s
197    /// module-level docs.
198    pub async fn resolve(&self, query: &Message) -> Result<Message> {
199        let correlation_id = self.next_correlation_id.fetch_add(1, Ordering::Relaxed);
200        let question = query.questions.first();
201        self.emit(ObserveEvent::QueryReceived {
202            correlation_id,
203            name: question.map(|question| question.name.clone()),
204            rtype: question.map(|question| question.qtype),
205            class: question.map(|question| question.qclass),
206        });
207        let Some(question) = query.questions.first() else {
208            self.emit(ObserveEvent::Failed {
209                correlation_id,
210                failure: ObserveFailure::NoRoute,
211            });
212            return Err(Error::NoRoute);
213        };
214        if let Some(fake_ip) = &self.fake_ip {
215            match fake_ip_answer(query, fake_ip) {
216                Ok(Some(answer)) => {
217                    self.emit(ObserveEvent::FakeIpTerminal { correlation_id });
218                    self.emit(ObserveEvent::Completed {
219                        correlation_id,
220                        rcode: answer.header.rcode,
221                    });
222                    return Ok(answer);
223                }
224                Ok(None) => {}
225                Err(error) => {
226                    self.emit(ObserveEvent::Failed {
227                        correlation_id,
228                        failure: observe_failure(&error),
229                    });
230                    return Err(error);
231                }
232            }
233        }
234
235        let (group, backends) = match self.select_backends(question, correlation_id).await {
236            Ok(selected) => selected,
237            Err(error) => {
238                self.emit(ObserveEvent::Failed {
239                    correlation_id,
240                    failure: observe_failure(&error),
241                });
242                return Err(error);
243            }
244        };
245        let key = CacheKey {
246            name: question.name.clone(),
247            rtype: question.qtype,
248            class: question.qclass,
249            group: group.clone(),
250        };
251
252        let now = self.clock.now();
253        {
254            let cache = self.cache.lock().expect("cache mutex poisoned");
255            if let Some(entry) = cache.get(&key)
256                && entry.expires_at > now
257            {
258                let answer = cache_hit_response(query, &entry.answer);
259                drop(cache);
260                self.emit(ObserveEvent::CacheHit {
261                    correlation_id,
262                    group: group.clone(),
263                });
264                self.emit(ObserveEvent::Completed {
265                    correlation_id,
266                    rcode: answer.header.rcode,
267                });
268                return Ok(answer);
269            }
270        }
271        self.emit(ObserveEvent::CacheMiss {
272            correlation_id,
273            group: group.clone(),
274        });
275
276        let mut last_err = None;
277        for (backend_index, backend) in backends.iter().enumerate() {
278            self.emit(ObserveEvent::UpstreamAttempt {
279                correlation_id,
280                group: group.clone(),
281                backend_index,
282            });
283            match backend.resolve(query).await {
284                Ok(answer) => {
285                    self.emit(ObserveEvent::UpstreamOutcome {
286                        correlation_id,
287                        group: group.clone(),
288                        backend_index,
289                        outcome: UpstreamObserveOutcome::Success,
290                    });
291                    if let Some(ttl) = cacheable_ttl(&answer) {
292                        let mut cache = self.cache.lock().expect("cache mutex poisoned");
293                        cache.insert(
294                            key,
295                            CacheEntry {
296                                answer: answer.clone(),
297                                expires_at: now + ttl,
298                            },
299                        );
300                    }
301                    self.emit(ObserveEvent::Completed {
302                        correlation_id,
303                        rcode: answer.header.rcode,
304                    });
305                    return Ok(answer);
306                }
307                Err(e) if is_retryable(&e) => {
308                    self.emit(ObserveEvent::UpstreamOutcome {
309                        correlation_id,
310                        group: group.clone(),
311                        backend_index,
312                        outcome: UpstreamObserveOutcome::RetryableFailure,
313                    });
314                    last_err = Some(e);
315                }
316                Err(e) => {
317                    self.emit(ObserveEvent::UpstreamOutcome {
318                        correlation_id,
319                        group: group.clone(),
320                        backend_index,
321                        outcome: UpstreamObserveOutcome::Failure,
322                    });
323                    self.emit(ObserveEvent::Failed {
324                        correlation_id,
325                        failure: observe_failure(&e),
326                    });
327                    return Err(e);
328                }
329            }
330        }
331
332        let error = last_err.expect("at least one backend was tried since backends is non-empty");
333        self.emit(ObserveEvent::Failed {
334            correlation_id,
335            failure: observe_failure(&error),
336        });
337        Err(error)
338    }
339
340    fn emit(&self, event: ObserveEvent) {
341        if let Some(sink) = &self.observability_sink {
342            let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| sink.record(&event)));
343        }
344    }
345
346    /// Selects and validates the effective upstream group for one ordinary
347    /// query. This deliberately happens before the cache lookup because a
348    /// hook may choose different groups for equal DNS questions.
349    ///
350    /// No resolver mutex is held while invoking the hook. Dropping the
351    /// enclosing [`Resolver::resolve`] future drops this in-flight hook call;
352    /// hook implementations own cancellation cleanup and must not re-enter
353    /// this resolver.
354    async fn select_backends(
355        &self,
356        question: &dns_lattice_model::Question,
357        correlation_id: u64,
358    ) -> Result<(UpstreamGroupId, &Vec<Box<dyn UpstreamBackend>>)> {
359        let static_group = self.policy.resolve_group(&question.name);
360        self.emit(ObserveEvent::StaticRoute {
361            correlation_id,
362            group: static_group.cloned(),
363        });
364        let group = match &self.route_hook {
365            Some(hook) => match hook.select(RouteRequest::new(question, static_group)).await {
366                Ok(RouteDecision::Use(group)) => {
367                    self.emit(ObserveEvent::HookDecision {
368                        correlation_id,
369                        decision: HookObserveDecision::Use(group.clone()),
370                    });
371                    Some(group)
372                }
373                Ok(RouteDecision::Abstain) => {
374                    self.emit(ObserveEvent::HookDecision {
375                        correlation_id,
376                        decision: HookObserveDecision::Abstain,
377                    });
378                    static_group.cloned()
379                }
380                Err(error) => {
381                    self.emit(ObserveEvent::HookDecision {
382                        correlation_id,
383                        decision: HookObserveDecision::Failed,
384                    });
385                    return Err(Error::Hook(error.to_string()));
386                }
387            },
388            None => static_group.cloned(),
389        }
390        .ok_or(Error::NoRoute)?;
391
392        let backends = self.backends.get(&group).ok_or(Error::NoRoute)?;
393        if backends.is_empty() {
394            return Err(Error::NoRoute);
395        }
396        Ok((group, backends))
397    }
398}
399
400/// Projects the transaction-specific parts of `query` onto a cached DNS
401/// response.
402///
403/// Cache entries represent reusable answer content, not a prior client's DNS
404/// transaction. The cache key intentionally covers only the first question's
405/// matching fields and effective upstream group, while a response must echo
406/// the current request's transaction ID and question section. All cached
407/// response flags and record sections remain unchanged.
408fn cache_hit_response(query: &Message, cached: &Message) -> Message {
409    let mut response = cached.clone();
410    response.header.id = query.header.id;
411    response.questions = query.questions.clone();
412    response
413}
414
415/// Returns a locally synthesized Fake IP response when `query` is handled by
416/// `fake_ip`, or `None` when the ordinary resolver pipeline must handle it.
417///
418/// Only IN A, IN AAAA, and canonical IN reverse PTR questions can be handled. A
419/// synthesized response intentionally bypasses the resolver cache and every
420/// upstream backend: its lifetime is the pool mapping lifetime and the pool
421/// is the authority for its reverse range.
422fn fake_ip_answer(query: &Message, fake_ip: &FakeIpResolverConfig) -> Result<Option<Message>> {
423    let Some(question) = query.questions.first() else {
424        return Ok(None);
425    };
426    if question.qclass != Class::In {
427        return Ok(None);
428    }
429
430    match question.qtype {
431        RecordType::A if fake_ip.policy.matches(&question.name) => {
432            if !fake_ip.pool.ipv4_enabled() {
433                return Ok(Some(local_response(query, Rcode::NoError)));
434            }
435            fake_ip_ttl(fake_ip.pool.ttl())?;
436            let mut answer = local_response(query, Rcode::NoError);
437            match fake_ip.pool.allocate_ipv4_with_ttl(question.name.clone()) {
438                Ok((address, lifetime)) => answer.answers.push(ResourceRecord {
439                    name: question.name.clone(),
440                    rtype: RecordType::A,
441                    class: Class::In,
442                    ttl: fake_ip_ttl(lifetime)?,
443                    rdata: RData::A(address),
444                }),
445                Err(Error::FakeIpFamilyDisabled) => {}
446                Err(error) => return Err(error),
447            }
448            Ok(Some(answer))
449        }
450        RecordType::Aaaa if fake_ip.policy.matches(&question.name) => {
451            if !fake_ip.pool.ipv6_enabled() {
452                return Ok(Some(local_response(query, Rcode::NoError)));
453            }
454            fake_ip_ttl(fake_ip.pool.ttl())?;
455            let mut answer = local_response(query, Rcode::NoError);
456            match fake_ip.pool.allocate_ipv6_with_ttl(question.name.clone()) {
457                Ok((address, lifetime)) => answer.answers.push(ResourceRecord {
458                    name: question.name.clone(),
459                    rtype: RecordType::Aaaa,
460                    class: Class::In,
461                    ttl: fake_ip_ttl(lifetime)?,
462                    rdata: RData::Aaaa(address),
463                }),
464                Err(Error::FakeIpFamilyDisabled) => {}
465                Err(error) => return Err(error),
466            }
467            Ok(Some(answer))
468        }
469        RecordType::Ptr => fake_ip_ptr_answer(query, fake_ip),
470        _ => Ok(None),
471    }
472}
473
474fn fake_ip_ptr_answer(query: &Message, fake_ip: &FakeIpResolverConfig) -> Result<Option<Message>> {
475    let question = query.questions.first().expect("checked by caller");
476    let address = match parse_reverse_name(&question.name) {
477        Some(address) => address,
478        None => return Ok(None),
479    };
480    let mapping = match address {
481        std::net::IpAddr::V4(address) if fake_ip.pool.contains_ipv4(address) => {
482            fake_ip.pool.lookup_ipv4_with_ttl(address)
483        }
484        std::net::IpAddr::V6(address) if fake_ip.pool.contains_ipv6(address) => {
485            fake_ip.pool.lookup_ipv6_with_ttl(address)
486        }
487        _ => return Ok(None),
488    };
489    let mut answer = local_response(
490        query,
491        if mapping.is_some() {
492            Rcode::NoError
493        } else {
494            Rcode::NxDomain
495        },
496    );
497    if let Some((name, lifetime)) = mapping {
498        answer.answers.push(ResourceRecord {
499            name: question.name.clone(),
500            rtype: RecordType::Ptr,
501            class: Class::In,
502            ttl: fake_ip_ttl(lifetime)?,
503            rdata: RData::Ptr(name),
504        });
505    }
506    Ok(Some(answer))
507}
508
509fn local_response(query: &Message, rcode: Rcode) -> Message {
510    let mut header = query.header;
511    header.qr = true;
512    header.rcode = rcode;
513    Message {
514        header,
515        questions: query.questions.clone(),
516        answers: Vec::new(),
517        authorities: Vec::new(),
518        additionals: Vec::new(),
519    }
520}
521
522fn fake_ip_ttl(lifetime: Duration) -> Result<u32> {
523    u32::try_from(lifetime.as_secs()).map_err(|_| Error::FakeIpTtlOutOfRange)
524}
525
526/// Parses a canonical `in-addr.arpa` or `ip6.arpa` owner name.
527///
528/// Non-canonical reverse names are intentionally routed normally, so only a
529/// pool range for which this resolver is authoritative receives local DNS
530/// semantics.
531fn parse_reverse_name(name: &Name) -> Option<std::net::IpAddr> {
532    let labels: Vec<_> = name.labels().collect();
533    if labels.len() == 6
534        && labels[4].eq_ignore_ascii_case(b"in-addr")
535        && labels[5].eq_ignore_ascii_case(b"arpa")
536    {
537        let mut octets = [0_u8; 4];
538        for (index, label) in labels[..4].iter().enumerate() {
539            let text = std::str::from_utf8(label).ok()?;
540            let value = text.parse::<u8>().ok()?;
541            if value.to_string() != text {
542                return None;
543            }
544            octets[3 - index] = value;
545        }
546        return Some(std::net::IpAddr::V4(Ipv4Addr::from(octets)));
547    }
548    if labels.len() == 34
549        && labels[32].eq_ignore_ascii_case(b"ip6")
550        && labels[33].eq_ignore_ascii_case(b"arpa")
551    {
552        let mut bytes = [0_u8; 16];
553        for (index, label) in labels[..32].iter().enumerate() {
554            if label.len() != 1 {
555                return None;
556            }
557            let nibble = match label[0] {
558                b'0'..=b'9' => label[0] - b'0',
559                b'a'..=b'f' => label[0] - b'a' + 10,
560                b'A'..=b'F' => label[0] - b'A' + 10,
561                _ => return None,
562            };
563            let target = 31 - index;
564            if target % 2 == 0 {
565                bytes[target / 2] |= nibble << 4;
566            } else {
567                bytes[target / 2] |= nibble;
568            }
569        }
570        return Some(std::net::IpAddr::V6(Ipv6Addr::from(bytes)));
571    }
572    None
573}
574
575/// Returns whether `err` should cause the failover loop to try the next
576/// backend in the group rather than propagate immediately: all three
577/// backend-level failure variants โ€”
578/// [`Error::Timeout`], [`Error::Transport`], and [`Error::Tls`] โ€” are
579/// retryable, since none indicate a client-input problem and a different
580/// backend in the same group may have independent connectivity/TLS
581/// configuration.
582fn is_retryable(err: &Error) -> bool {
583    matches!(err, Error::Timeout | Error::Transport(_) | Error::Tls(_))
584}
585
586fn observe_failure(error: &Error) -> ObserveFailure {
587    match error {
588        Error::NoRoute => ObserveFailure::NoRoute,
589        Error::Hook(_) => ObserveFailure::Hook,
590        Error::Timeout => ObserveFailure::Timeout,
591        Error::Transport(_) => ObserveFailure::Transport,
592        Error::Tls(_) => ObserveFailure::Tls,
593        _ => ObserveFailure::Other,
594    }
595}
596
597/// Determines the [`Duration`] an `answer` should be cached for, or `None`
598/// if it should not be cached at all.
599///
600/// Positive answers (`NoError` with at least one answer record) use the
601/// minimum `ttl` across their answer records. Negative
602/// answers (`NxDomain`, or `NoError` with an empty answer section) use the
603/// `minimum` field of an SOA record in the authority section when present
604/// (RFC 2308 ยง5), else [`NEGATIVE_CACHE_FLOOR`].
605fn cacheable_ttl(answer: &Message) -> Option<Duration> {
606    let is_negative = matches!(answer.header.rcode, Rcode::NxDomain)
607        || (matches!(answer.header.rcode, Rcode::NoError) && answer.answers.is_empty());
608
609    if is_negative {
610        let ttl = answer
611            .authorities
612            .iter()
613            .find_map(|rr| match &rr.rdata {
614                RData::Soa { minimum, .. } => Some(*minimum),
615                _ => None,
616            })
617            .map(|minimum| Duration::from_secs(u64::from(minimum)))
618            .unwrap_or(NEGATIVE_CACHE_FLOOR);
619        return Some(ttl);
620    }
621
622    if answer.answers.is_empty() {
623        return None;
624    }
625
626    answer
627        .answers
628        .iter()
629        .map(|rr| rr.ttl)
630        .min()
631        .map(|ttl| Duration::from_secs(u64::from(ttl)))
632}
633
634/// Builds a [`Resolver`] from a split-DNS policy and one or more upstream
635/// backends per group.
636pub struct ResolverBuilder {
637    policy: SplitDnsPolicy,
638    backends: HashMap<UpstreamGroupId, Vec<Box<dyn UpstreamBackend>>>,
639    clock: Box<dyn Clock + Send + Sync>,
640    fake_ip: Option<FakeIpResolverConfig>,
641    route_hook: Option<Box<dyn RouteHook>>,
642    observability_sink: Option<Arc<dyn ObservabilitySink>>,
643}
644
645impl ResolverBuilder {
646    /// Registers an upstream backend used to answer queries routed to
647    /// `group`, appended after any backend already registered for that
648    /// group. [`Resolver::resolve`] tries a group's backends in this
649    /// registration order, falling over to the next one on a retryable
650    /// error.
651    ///
652    /// `backend` is any [`crate::upstream::UpstreamBackend`] implementation
653    /// โ€” e.g. [`crate::upstream::UdpBackend`]/
654    /// [`crate::upstream::TcpBackend`] for real transport, or a
655    /// test-only fake implementing the trait directly.
656    pub fn backend(
657        mut self,
658        group: UpstreamGroupId,
659        backend: impl UpstreamBackend + 'static,
660    ) -> Self {
661        self.backends
662            .entry(group)
663            .or_default()
664            .push(Box::new(backend));
665        self
666    }
667
668    /// Enables local Fake IP synthesis for names selected by `policy`.
669    ///
670    /// Matching IN A/AAAA questions allocate or reuse an address in `pool`
671    /// and return a local response without consulting the cache or an
672    /// upstream. If the selected address family is disabled in `pool`, the
673    /// resolver instead returns a local NOERROR empty answer (NODATA), still
674    /// without a cache or upstream lookup. Canonical IN PTR questions inside
675    /// one of the pool's ranges are likewise handled locally: live mappings
676    /// return PTR, and an unmapped address returns NXDOMAIN. All other
677    /// questions follow normal split-DNS resolution.
678    pub fn fake_ip(mut self, pool: Arc<FakeIpPool>, policy: FakeIpPolicy) -> Self {
679        self.fake_ip = Some(FakeIpResolverConfig { pool, policy });
680        self
681    }
682
683    /// Configures one optional dynamic upstream-group selection hook.
684    ///
685    /// For each non-local query, the hook receives the first question and
686    /// the static split-DNS candidate. [`crate::hooks::RouteDecision::Use`]
687    /// replaces that candidate, while `Abstain` retains it. The resulting
688    /// group must be registered and nonempty; otherwise resolution returns
689    /// [`Error::NoRoute`] without static fallback, cache access, or an
690    /// upstream call. Fake IP local answers remain terminal and never invoke
691    /// the hook.
692    ///
693    /// The hook owns timeout, retry, and cancellation cleanup. A dropped
694    /// [`Resolver::resolve`] call drops the in-flight hook future. Hooks must
695    /// not call `resolve` on this same resolver directly or indirectly.
696    pub fn route_hook(mut self, hook: impl RouteHook + 'static) -> Self {
697        self.route_hook = Some(Box::new(hook));
698        self
699    }
700
701    /// Adds an optional, non-authoritative synchronous event sink.
702    ///
703    /// The resolver invokes it outside internal locks. Panics are isolated and
704    /// never affect DNS answers, route selection, cache behavior, or retries.
705    pub fn observability_sink(mut self, sink: Arc<dyn ObservabilitySink>) -> Self {
706        self.observability_sink = Some(sink);
707        self
708    }
709
710    /// Substitutes the clock used to compute and check cache expiry.
711    /// Crate-private: no public API for clock injection.
712    #[cfg(test)]
713    pub(crate) fn clock(mut self, clock: impl Clock + Send + Sync + 'static) -> Self {
714        self.clock = Box::new(clock);
715        self
716    }
717
718    /// Builds the resolver.
719    pub fn build(self) -> Resolver {
720        Resolver {
721            policy: self.policy,
722            backends: self.backends,
723            clock: self.clock,
724            cache: Mutex::new(HashMap::new()),
725            fake_ip: self.fake_ip,
726            route_hook: self.route_hook,
727            observability_sink: self.observability_sink,
728            next_correlation_id: AtomicU64::new(1),
729        }
730    }
731}
732
733#[cfg(test)]
734mod tests {
735    use super::*;
736    use async_trait::async_trait;
737    use dns_lattice_model::{
738        Class, DomainPattern, Header, Name, Opcode, Question, Rcode, RecordType,
739    };
740
741    /// A minimal test-only fake [`UpstreamBackend`] that always returns a
742    /// fixed answer, proving routing wiring without modelling any real
743    /// upstream transport behavior.
744    struct FixedBackend(Message);
745
746    #[async_trait]
747    impl UpstreamBackend for FixedBackend {
748        async fn resolve(&self, _query: &Message) -> Result<Message> {
749            Ok(self.0.clone())
750        }
751    }
752
753    fn fixed_backend(answer: Message) -> FixedBackend {
754        FixedBackend(answer)
755    }
756
757    /// A test-only fake [`UpstreamBackend`] that always fails with a fixed
758    /// error, proving error propagation without modelling any real
759    /// upstream transport behavior.
760    struct FailingBackend(Error);
761
762    #[async_trait]
763    impl UpstreamBackend for FailingBackend {
764        async fn resolve(&self, _query: &Message) -> Result<Message> {
765            Err(self.0.clone())
766        }
767    }
768
769    fn n(s: &str) -> Name {
770        Name::from_ascii(s).unwrap()
771    }
772
773    fn query_for(name: &str) -> Message {
774        Message {
775            header: Header {
776                id: 1,
777                qr: false,
778                opcode: Opcode::Query,
779                authoritative: false,
780                truncated: false,
781                recursion_desired: true,
782                recursion_available: false,
783                rcode: Rcode::NoError,
784            },
785            questions: vec![Question {
786                name: n(name),
787                qtype: RecordType::A,
788                qclass: Class::In,
789            }],
790            answers: vec![],
791            authorities: vec![],
792            additionals: vec![],
793        }
794    }
795
796    fn answer_tagged(id: u16) -> Message {
797        let mut msg = query_for("tag.example");
798        msg.header.id = id;
799        msg.header.qr = true;
800        msg
801    }
802
803    #[tokio::test]
804    async fn routes_exact_match_to_its_group() {
805        let policy = SplitDnsPolicy::builder()
806            .rule(
807                DomainPattern::exact(n("host.corp.internal")),
808                UpstreamGroupId::new("corp"),
809            )
810            .build();
811        let resolver = Resolver::builder(policy)
812            .backend(
813                UpstreamGroupId::new("corp"),
814                fixed_backend(answer_tagged(42)),
815            )
816            .build();
817
818        let answer = resolver
819            .resolve(&query_for("host.corp.internal"))
820            .await
821            .expect("routed to corp backend");
822        assert_eq!(answer.header.id, 42);
823    }
824
825    #[tokio::test]
826    async fn routes_suffix_match_to_its_group() {
827        let policy = SplitDnsPolicy::builder()
828            .rule(
829                DomainPattern::suffix(n("corp.internal")),
830                UpstreamGroupId::new("corp"),
831            )
832            .build();
833        let resolver = Resolver::builder(policy)
834            .backend(
835                UpstreamGroupId::new("corp"),
836                fixed_backend(answer_tagged(7)),
837            )
838            .build();
839
840        let answer = resolver
841            .resolve(&query_for("host.corp.internal"))
842            .await
843            .expect("routed to corp backend via suffix");
844        assert_eq!(answer.header.id, 7);
845    }
846
847    #[tokio::test]
848    async fn routes_wildcard_match_to_its_group() {
849        let policy = SplitDnsPolicy::builder()
850            .rule(
851                DomainPattern::wildcard(n("corp.internal")),
852                UpstreamGroupId::new("corp"),
853            )
854            .build();
855        let resolver = Resolver::builder(policy)
856            .backend(
857                UpstreamGroupId::new("corp"),
858                fixed_backend(answer_tagged(9)),
859            )
860            .build();
861
862        let answer = resolver
863            .resolve(&query_for("host.corp.internal"))
864            .await
865            .expect("routed to corp backend via wildcard");
866        assert_eq!(answer.header.id, 9);
867    }
868
869    #[tokio::test]
870    async fn routes_unmatched_query_to_default_group() {
871        let policy = SplitDnsPolicy::builder()
872            .rule(
873                DomainPattern::suffix(n("corp.internal")),
874                UpstreamGroupId::new("corp"),
875            )
876            .default_group(UpstreamGroupId::new("public"))
877            .build();
878        let resolver = Resolver::builder(policy)
879            .backend(
880                UpstreamGroupId::new("public"),
881                fixed_backend(answer_tagged(3)),
882            )
883            .build();
884
885        let answer = resolver
886            .resolve(&query_for("example.com"))
887            .await
888            .expect("routed to default group");
889        assert_eq!(answer.header.id, 3);
890    }
891
892    #[tokio::test]
893    async fn no_route_when_no_match_and_no_default_group() {
894        let policy = SplitDnsPolicy::builder().build();
895        let resolver = Resolver::builder(policy).build();
896
897        let err = resolver
898            .resolve(&query_for("example.com"))
899            .await
900            .expect_err("no rule and no default group configured");
901        assert_eq!(err, Error::NoRoute);
902    }
903
904    #[tokio::test]
905    async fn no_route_when_matched_group_has_no_registered_backend() {
906        let policy = SplitDnsPolicy::builder()
907            .rule(
908                DomainPattern::suffix(n("corp.internal")),
909                UpstreamGroupId::new("corp"),
910            )
911            .build();
912        let resolver = Resolver::builder(policy).build();
913
914        let err = resolver
915            .resolve(&query_for("host.corp.internal"))
916            .await
917            .expect_err("matched group has no backend registered");
918        assert_eq!(err, Error::NoRoute);
919    }
920
921    #[tokio::test]
922    async fn failover_first_backend_succeeds_second_never_called() {
923        let policy = SplitDnsPolicy::builder()
924            .default_group(UpstreamGroupId::new("g"))
925            .build();
926        let calls = Arc::new(AtomicUsize::new(0));
927        let second = CountingBackend {
928            answer: answer_tagged(2),
929            calls: calls.clone(),
930        };
931        let resolver = Resolver::builder(policy)
932            .backend(UpstreamGroupId::new("g"), fixed_backend(answer_tagged(1)))
933            .backend(UpstreamGroupId::new("g"), second)
934            .build();
935
936        let answer = resolver
937            .resolve(&query_for("example.com"))
938            .await
939            .expect("first backend answers");
940        assert_eq!(answer.header.id, 1);
941        assert_eq!(
942            calls.load(Ordering::SeqCst),
943            0,
944            "second backend never called once the first succeeds"
945        );
946    }
947
948    #[tokio::test]
949    async fn failover_first_backend_fails_second_succeeds() {
950        let policy = SplitDnsPolicy::builder()
951            .default_group(UpstreamGroupId::new("g"))
952            .build();
953        let resolver = Resolver::builder(policy)
954            .backend(UpstreamGroupId::new("g"), FailingBackend(Error::Timeout))
955            .backend(UpstreamGroupId::new("g"), fixed_backend(answer_tagged(99)))
956            .build();
957
958        let answer = resolver
959            .resolve(&query_for("example.com"))
960            .await
961            .expect("second backend answers after first times out");
962        assert_eq!(
963            answer.header.id, 99,
964            "routed answer is the second backend's"
965        );
966    }
967
968    #[tokio::test]
969    async fn failover_tls_error_retries_to_next_backend() {
970        let policy = SplitDnsPolicy::builder()
971            .default_group(UpstreamGroupId::new("g"))
972            .build();
973        let resolver = Resolver::builder(policy)
974            .backend(
975                UpstreamGroupId::new("g"),
976                FailingBackend(Error::Tls("certificate expired".to_string())),
977            )
978            .backend(UpstreamGroupId::new("g"), fixed_backend(answer_tagged(5)))
979            .build();
980
981        let answer = resolver
982            .resolve(&query_for("example.com"))
983            .await
984            .expect("tls error on first backend retries to the second");
985        assert_eq!(answer.header.id, 5);
986    }
987
988    #[tokio::test]
989    async fn failover_all_backends_fail_returns_last_error() {
990        let policy = SplitDnsPolicy::builder()
991            .default_group(UpstreamGroupId::new("g"))
992            .build();
993        let resolver = Resolver::builder(policy)
994            .backend(UpstreamGroupId::new("g"), FailingBackend(Error::Timeout))
995            .backend(
996                UpstreamGroupId::new("g"),
997                FailingBackend(Error::Transport("connection refused".to_string())),
998            )
999            .build();
1000
1001        let err = resolver
1002            .resolve(&query_for("example.com"))
1003            .await
1004            .expect_err("both backends fail");
1005        assert_eq!(
1006            err,
1007            Error::Transport("connection refused".to_string()),
1008            "the last attempted backend's error is returned, not the first's"
1009        );
1010    }
1011
1012    #[tokio::test]
1013    async fn single_backend_group_still_behaves_as_before() {
1014        let policy = SplitDnsPolicy::builder()
1015            .default_group(UpstreamGroupId::new("g"))
1016            .build();
1017        let resolver = Resolver::builder(policy)
1018            .backend(UpstreamGroupId::new("g"), fixed_backend(answer_tagged(11)))
1019            .build();
1020
1021        let answer = resolver
1022            .resolve(&query_for("example.com"))
1023            .await
1024            .expect("single-backend group still resolves");
1025        assert_eq!(answer.header.id, 11);
1026    }
1027
1028    #[tokio::test]
1029    async fn backend_error_propagates_as_is() {
1030        let policy = SplitDnsPolicy::builder()
1031            .rule(
1032                DomainPattern::suffix(n("corp.internal")),
1033                UpstreamGroupId::new("corp"),
1034            )
1035            .build();
1036        let resolver = Resolver::builder(policy)
1037            .backend(
1038                UpstreamGroupId::new("corp"),
1039                FailingBackend(Error::NameTooLong),
1040            )
1041            .build();
1042
1043        let err = resolver
1044            .resolve(&query_for("host.corp.internal"))
1045            .await
1046            .expect_err("backend failure propagates");
1047        assert_eq!(err, Error::NameTooLong);
1048    }
1049
1050    // --- Dedicated fake upstream backend + cache test suite (deferred from
1051    // the routing and cache slice -----------------------------------------
1052
1053    use std::net::Ipv4Addr;
1054    use std::sync::Arc;
1055    use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
1056
1057    use dns_lattice_model::{RData, ResourceRecord};
1058    use tokio::sync::Notify;
1059
1060    use crate::hooks::{RouteDecision, RouteHook, RouteHookError, RouteRequest};
1061    use crate::observability::{ObservabilitySink, ObserveEvent};
1062
1063    #[derive(Clone)]
1064    struct PoolClock(Arc<Mutex<Instant>>);
1065
1066    impl PoolClock {
1067        fn new() -> Self {
1068            Self(Arc::new(Mutex::new(Instant::now())))
1069        }
1070
1071        fn advance(&self, duration: Duration) {
1072            *self.0.lock().expect("pool clock mutex poisoned") += duration;
1073        }
1074    }
1075
1076    impl crate::fakeip::Clock for PoolClock {
1077        fn now(&self) -> Instant {
1078            *self.0.lock().expect("pool clock mutex poisoned")
1079        }
1080    }
1081
1082    /// A configurable fake in-process [`UpstreamBackend`] that returns a
1083    /// fixed answer and counts how many times it was called, so cache-hit
1084    /// tests can assert the backend is *not* called again on a hit.
1085    struct CountingBackend {
1086        answer: Message,
1087        calls: Arc<AtomicUsize>,
1088    }
1089
1090    #[derive(Default)]
1091    struct RecordingSink(Mutex<Vec<ObserveEvent>>);
1092
1093    impl ObservabilitySink for RecordingSink {
1094        fn record(&self, event: &ObserveEvent) {
1095            self.0
1096                .lock()
1097                .expect("event mutex poisoned")
1098                .push(event.clone());
1099        }
1100    }
1101
1102    struct PanickingSink;
1103
1104    impl ObservabilitySink for PanickingSink {
1105        fn record(&self, _: &ObserveEvent) {
1106            panic!("observer failure must be isolated");
1107        }
1108    }
1109
1110    #[async_trait]
1111    impl UpstreamBackend for CountingBackend {
1112        async fn resolve(&self, _query: &Message) -> Result<Message> {
1113            self.calls.fetch_add(1, Ordering::SeqCst);
1114            Ok(self.answer.clone())
1115        }
1116    }
1117
1118    struct FixedHook {
1119        decision: std::result::Result<RouteDecision, RouteHookError>,
1120        calls: Arc<AtomicUsize>,
1121    }
1122
1123    #[async_trait]
1124    impl RouteHook for FixedHook {
1125        async fn select(
1126            &self,
1127            _request: RouteRequest<'_>,
1128        ) -> std::result::Result<RouteDecision, RouteHookError> {
1129            self.calls.fetch_add(1, Ordering::SeqCst);
1130            self.decision.clone()
1131        }
1132    }
1133
1134    struct SequencedHook {
1135        decisions: Mutex<Vec<RouteDecision>>,
1136    }
1137
1138    #[async_trait]
1139    impl RouteHook for SequencedHook {
1140        async fn select(
1141            &self,
1142            _request: RouteRequest<'_>,
1143        ) -> std::result::Result<RouteDecision, RouteHookError> {
1144            Ok(self
1145                .decisions
1146                .lock()
1147                .expect("hook decisions mutex poisoned")
1148                .remove(0))
1149        }
1150    }
1151
1152    struct RecordingHook {
1153        decision: RouteDecision,
1154        static_groups: Arc<Mutex<Vec<Option<UpstreamGroupId>>>>,
1155    }
1156
1157    #[async_trait]
1158    impl RouteHook for RecordingHook {
1159        async fn select(
1160            &self,
1161            request: RouteRequest<'_>,
1162        ) -> std::result::Result<RouteDecision, RouteHookError> {
1163            self.static_groups
1164                .lock()
1165                .expect("recorded static groups mutex poisoned")
1166                .push(request.static_group().cloned());
1167            Ok(self.decision.clone())
1168        }
1169    }
1170
1171    struct PendingHook {
1172        entered: Arc<Notify>,
1173        dropped: Arc<AtomicBool>,
1174    }
1175
1176    struct DropSignal(Arc<AtomicBool>);
1177
1178    impl Drop for DropSignal {
1179        fn drop(&mut self) {
1180            self.0.store(true, Ordering::SeqCst);
1181        }
1182    }
1183
1184    #[async_trait]
1185    impl RouteHook for PendingHook {
1186        async fn select(
1187            &self,
1188            _request: RouteRequest<'_>,
1189        ) -> std::result::Result<RouteDecision, RouteHookError> {
1190            let _drop_signal = DropSignal(self.dropped.clone());
1191            self.entered.notify_waiters();
1192            std::future::pending().await
1193        }
1194    }
1195
1196    fn a_answer(name: &str, ttl: u32) -> Message {
1197        let mut msg = query_for(name);
1198        msg.header.qr = true;
1199        msg.answers.push(ResourceRecord {
1200            name: n(name),
1201            rtype: RecordType::A,
1202            class: Class::In,
1203            ttl,
1204            rdata: RData::A(Ipv4Addr::new(203, 0, 113, 1)),
1205        });
1206        msg
1207    }
1208
1209    fn nxdomain_answer(name: &str, soa_minimum: Option<u32>) -> Message {
1210        let mut msg = query_for(name);
1211        msg.header.qr = true;
1212        msg.header.rcode = Rcode::NxDomain;
1213        if let Some(minimum) = soa_minimum {
1214            msg.authorities.push(ResourceRecord {
1215                name: n("example.com"),
1216                rtype: RecordType::Soa,
1217                class: Class::In,
1218                ttl: 3600,
1219                rdata: RData::Soa {
1220                    mname: n("ns1.example.com"),
1221                    rname: n("hostmaster.example.com"),
1222                    serial: 1,
1223                    refresh: 3600,
1224                    retry: 600,
1225                    expire: 604_800,
1226                    minimum,
1227                },
1228            });
1229        }
1230        msg
1231    }
1232
1233    fn nodata_answer(name: &str) -> Message {
1234        // NoError, empty answer section: NODATA per RFC 2308.
1235        query_for_response(name)
1236    }
1237
1238    fn query_for_response(name: &str) -> Message {
1239        let mut msg = query_for(name);
1240        msg.header.qr = true;
1241        msg
1242    }
1243
1244    fn resolver_with_counting_backend(
1245        policy: SplitDnsPolicy,
1246        group: &str,
1247        answer: Message,
1248        clock: FakeClock,
1249    ) -> (Resolver, Arc<AtomicUsize>) {
1250        let calls = Arc::new(AtomicUsize::new(0));
1251        let backend = CountingBackend {
1252            answer,
1253            calls: calls.clone(),
1254        };
1255        let resolver = Resolver::builder(policy)
1256            .clock(clock)
1257            .backend(UpstreamGroupId::new(group), backend)
1258            .build();
1259        (resolver, calls)
1260    }
1261
1262    #[tokio::test]
1263    async fn cache_hit_does_not_call_backend_again() {
1264        let policy = SplitDnsPolicy::builder()
1265            .default_group(UpstreamGroupId::new("g"))
1266            .build();
1267        let (resolver, calls) = resolver_with_counting_backend(
1268            policy,
1269            "g",
1270            a_answer("example.com", 300),
1271            FakeClock::new(),
1272        );
1273
1274        let first = resolver
1275            .resolve(&query_for("example.com"))
1276            .await
1277            .expect("first resolve populates cache");
1278        let second = resolver
1279            .resolve(&query_for("example.com"))
1280            .await
1281            .expect("second resolve served from cache");
1282
1283        assert_eq!(first, second);
1284        assert_eq!(calls.load(Ordering::SeqCst), 1, "backend called only once");
1285    }
1286
1287    #[tokio::test]
1288    async fn observability_reports_ordered_cache_miss_and_hit_without_affecting_resolution() {
1289        let sink = Arc::new(RecordingSink::default());
1290        let policy = SplitDnsPolicy::builder()
1291            .default_group(UpstreamGroupId::new("g"))
1292            .build();
1293        let (base, calls) = resolver_with_counting_backend(
1294            policy,
1295            "g",
1296            a_answer("example.com", 300),
1297            FakeClock::new(),
1298        );
1299        let resolver = ResolverBuilder {
1300            policy: base.policy,
1301            backends: base.backends,
1302            clock: base.clock,
1303            fake_ip: base.fake_ip,
1304            route_hook: base.route_hook,
1305            observability_sink: Some(sink.clone()),
1306        }
1307        .build();
1308
1309        resolver.resolve(&query_for("example.com")).await.unwrap();
1310        resolver.resolve(&query_for("example.com")).await.unwrap();
1311        assert_eq!(calls.load(Ordering::SeqCst), 1);
1312
1313        let events = sink.0.lock().unwrap().clone();
1314        assert!(matches!(
1315            events[0],
1316            ObserveEvent::QueryReceived {
1317                correlation_id: 1,
1318                ..
1319            }
1320        ));
1321        assert!(matches!(
1322            events[1],
1323            ObserveEvent::StaticRoute {
1324                correlation_id: 1,
1325                ..
1326            }
1327        ));
1328        assert!(matches!(
1329            events[2],
1330            ObserveEvent::CacheMiss {
1331                correlation_id: 1,
1332                ..
1333            }
1334        ));
1335        assert!(matches!(
1336            events[3],
1337            ObserveEvent::UpstreamAttempt {
1338                correlation_id: 1,
1339                backend_index: 0,
1340                ..
1341            }
1342        ));
1343        assert!(matches!(
1344            events[4],
1345            ObserveEvent::UpstreamOutcome {
1346                correlation_id: 1,
1347                outcome: UpstreamObserveOutcome::Success,
1348                ..
1349            }
1350        ));
1351        assert!(matches!(
1352            events[5],
1353            ObserveEvent::Completed {
1354                correlation_id: 1,
1355                ..
1356            }
1357        ));
1358        assert!(matches!(
1359            events[6],
1360            ObserveEvent::QueryReceived {
1361                correlation_id: 2,
1362                ..
1363            }
1364        ));
1365        assert!(matches!(
1366            events[7],
1367            ObserveEvent::StaticRoute {
1368                correlation_id: 2,
1369                ..
1370            }
1371        ));
1372        assert!(matches!(
1373            events[8],
1374            ObserveEvent::CacheHit {
1375                correlation_id: 2,
1376                ..
1377            }
1378        ));
1379        assert!(matches!(
1380            events[9],
1381            ObserveEvent::Completed {
1382                correlation_id: 2,
1383                ..
1384            }
1385        ));
1386    }
1387
1388    #[tokio::test]
1389    async fn panicking_observability_sink_is_non_authoritative() {
1390        let resolver = Resolver::builder(
1391            SplitDnsPolicy::builder()
1392                .default_group(UpstreamGroupId::new("g"))
1393                .build(),
1394        )
1395        .backend(
1396            UpstreamGroupId::new("g"),
1397            fixed_backend(a_answer("example.com", 300)),
1398        )
1399        .observability_sink(Arc::new(PanickingSink))
1400        .build();
1401
1402        assert!(resolver.resolve(&query_for("example.com")).await.is_ok());
1403    }
1404
1405    #[tokio::test]
1406    async fn observability_starts_empty_queries_before_no_route_failure() {
1407        let sink = Arc::new(RecordingSink::default());
1408        let resolver = Resolver::builder(SplitDnsPolicy::builder().build())
1409            .observability_sink(sink.clone())
1410            .build();
1411        let mut query = query_for("example.com");
1412        query.questions.clear();
1413
1414        assert_eq!(resolver.resolve(&query).await, Err(Error::NoRoute));
1415        assert!(matches!(
1416            sink.0.lock().unwrap().as_slice(),
1417            [
1418                ObserveEvent::QueryReceived {
1419                    name: None,
1420                    rtype: None,
1421                    class: None,
1422                    ..
1423                },
1424                ObserveEvent::Failed {
1425                    failure: ObserveFailure::NoRoute,
1426                    ..
1427                },
1428            ]
1429        ));
1430    }
1431
1432    #[tokio::test]
1433    async fn observability_marks_fake_ip_terminal_before_cache_or_upstream() {
1434        let sink = Arc::new(RecordingSink::default());
1435        let resolver = Resolver::builder(SplitDnsPolicy::builder().build())
1436            .fake_ip(
1437                fake_ip_pool(PoolClock::new()),
1438                fake_ip_policy("example.com"),
1439            )
1440            .observability_sink(sink.clone())
1441            .build();
1442
1443        resolver.resolve(&query_for("example.com")).await.unwrap();
1444        assert!(matches!(
1445            sink.0.lock().unwrap().as_slice(),
1446            [
1447                ObserveEvent::QueryReceived { .. },
1448                ObserveEvent::FakeIpTerminal { .. },
1449                ObserveEvent::Completed { .. },
1450            ]
1451        ));
1452    }
1453
1454    #[tokio::test]
1455    async fn observability_records_hook_and_timeout_failures_in_order() {
1456        let sink = Arc::new(RecordingSink::default());
1457        let resolver = Resolver::builder(
1458            SplitDnsPolicy::builder()
1459                .default_group(UpstreamGroupId::new("g"))
1460                .build(),
1461        )
1462        .route_hook(FixedHook {
1463            decision: Err(RouteHookError::new("denied")),
1464            calls: Arc::new(AtomicUsize::new(0)),
1465        })
1466        .observability_sink(sink.clone())
1467        .build();
1468        assert!(matches!(
1469            resolver.resolve(&query_for("example.com")).await,
1470            Err(Error::Hook(_))
1471        ));
1472        assert!(matches!(
1473            sink.0.lock().unwrap().as_slice(),
1474            [
1475                ObserveEvent::QueryReceived { .. },
1476                ObserveEvent::StaticRoute { .. },
1477                ObserveEvent::HookDecision {
1478                    decision: HookObserveDecision::Failed,
1479                    ..
1480                },
1481                ObserveEvent::Failed {
1482                    failure: ObserveFailure::Hook,
1483                    ..
1484                },
1485            ]
1486        ));
1487
1488        sink.0.lock().unwrap().clear();
1489        let resolver = Resolver::builder(
1490            SplitDnsPolicy::builder()
1491                .default_group(UpstreamGroupId::new("g"))
1492                .build(),
1493        )
1494        .backend(UpstreamGroupId::new("g"), FailingBackend(Error::Timeout))
1495        .observability_sink(sink.clone())
1496        .build();
1497        assert_eq!(
1498            resolver.resolve(&query_for("example.com")).await,
1499            Err(Error::Timeout)
1500        );
1501        assert!(matches!(
1502            sink.0.lock().unwrap().as_slice(),
1503            [
1504                ObserveEvent::QueryReceived { .. },
1505                ObserveEvent::StaticRoute { .. },
1506                ObserveEvent::CacheMiss { .. },
1507                ObserveEvent::UpstreamAttempt { .. },
1508                ObserveEvent::UpstreamOutcome {
1509                    outcome: UpstreamObserveOutcome::RetryableFailure,
1510                    ..
1511                },
1512                ObserveEvent::Failed {
1513                    failure: ObserveFailure::Timeout,
1514                    ..
1515                },
1516            ]
1517        ));
1518    }
1519
1520    #[tokio::test]
1521    async fn cache_hit_preserves_the_current_query_identity_and_questions() {
1522        let policy = SplitDnsPolicy::builder()
1523            .default_group(UpstreamGroupId::new("g"))
1524            .build();
1525        let (resolver, calls) = resolver_with_counting_backend(
1526            policy,
1527            "g",
1528            a_answer("example.com", 300),
1529            FakeClock::new(),
1530        );
1531        let first = query_for_type("example.com", RecordType::A, Class::In, 91);
1532        let mut second = query_for_type("example.com", RecordType::A, Class::In, 92);
1533        second.questions.push(Question {
1534            name: n("extra.example.com"),
1535            qtype: RecordType::Aaaa,
1536            qclass: Class::In,
1537        });
1538
1539        resolver
1540            .resolve(&first)
1541            .await
1542            .expect("first resolve populates cache");
1543        let cached = resolver
1544            .resolve(&second)
1545            .await
1546            .expect("second resolve is served from cache");
1547
1548        assert_eq!(cached.header.id, 92);
1549        assert_eq!(cached.questions, second.questions);
1550        assert_eq!(
1551            cached.answers[0].rdata,
1552            RData::A(Ipv4Addr::new(203, 0, 113, 1))
1553        );
1554        assert_eq!(
1555            calls.load(Ordering::SeqCst),
1556            1,
1557            "second query is a cache hit"
1558        );
1559    }
1560
1561    #[tokio::test]
1562    async fn cache_identity_separates_generated_question_type_and_class_pairs() {
1563        let policy = SplitDnsPolicy::builder()
1564            .default_group(UpstreamGroupId::new("g"))
1565            .build();
1566        let (resolver, calls) = resolver_with_counting_backend(
1567            policy,
1568            "g",
1569            a_answer("example.com", 300),
1570            FakeClock::new(),
1571        );
1572
1573        let cases = [
1574            (RecordType::A, Class::In, 1),
1575            (RecordType::Aaaa, Class::In, 2),
1576            (RecordType::A, Class::Ch, 3),
1577            (RecordType::Other(65280), Class::Other(65280), 4),
1578        ];
1579
1580        for (rtype, class, id) in cases {
1581            resolver
1582                .resolve(&query_for_type("example.com", rtype, class, id))
1583                .await
1584                .expect("each distinct cache identity resolves");
1585        }
1586        assert_eq!(calls.load(Ordering::SeqCst), cases.len());
1587
1588        for (rtype, class, id) in cases {
1589            let cached = resolver
1590                .resolve(&query_for_type("example.com", rtype, class, id + 10))
1591                .await
1592                .expect("same type/class pair is cached");
1593            assert_eq!(cached.header.id, id + 10);
1594            assert_eq!(cached.questions[0].qtype, rtype);
1595            assert_eq!(cached.questions[0].qclass, class);
1596        }
1597        assert_eq!(calls.load(Ordering::SeqCst), cases.len());
1598    }
1599
1600    #[tokio::test]
1601    async fn cache_entry_still_hit_just_before_ttl_elapses() {
1602        let policy = SplitDnsPolicy::builder()
1603            .default_group(UpstreamGroupId::new("g"))
1604            .build();
1605        let clock = FakeClock::new();
1606        let (resolver, calls) = resolver_with_counting_backend(
1607            policy,
1608            "g",
1609            a_answer("example.com", 300),
1610            clock.clone(),
1611        );
1612
1613        resolver
1614            .resolve(&query_for("example.com"))
1615            .await
1616            .expect("first resolve populates cache");
1617        assert_eq!(calls.load(Ordering::SeqCst), 1);
1618
1619        clock.advance(Duration::from_secs(299));
1620
1621        resolver
1622            .resolve(&query_for("example.com"))
1623            .await
1624            .expect("still cached before ttl elapses");
1625        assert_eq!(calls.load(Ordering::SeqCst), 1, "cache hit before expiry");
1626    }
1627
1628    #[tokio::test]
1629    async fn negative_answer_is_cached_with_soa_minimum_ttl() {
1630        let policy = SplitDnsPolicy::builder()
1631            .default_group(UpstreamGroupId::new("g"))
1632            .build();
1633        let (resolver, calls) = resolver_with_counting_backend(
1634            policy,
1635            "g",
1636            nxdomain_answer("missing.example.com", Some(300)),
1637            FakeClock::new(),
1638        );
1639
1640        let first = resolver
1641            .resolve(&query_for("missing.example.com"))
1642            .await
1643            .expect("nxdomain is Ok(Message), not Err");
1644        assert_eq!(first.header.rcode, Rcode::NxDomain);
1645        resolver
1646            .resolve(&query_for("missing.example.com"))
1647            .await
1648            .expect("served from negative cache");
1649        assert_eq!(calls.load(Ordering::SeqCst), 1, "negative answer cached");
1650    }
1651
1652    #[tokio::test]
1653    async fn negative_answer_without_soa_uses_fixed_floor_ttl() {
1654        let policy = SplitDnsPolicy::builder()
1655            .default_group(UpstreamGroupId::new("g"))
1656            .build();
1657        let (resolver, calls) = resolver_with_counting_backend(
1658            policy,
1659            "g",
1660            nxdomain_answer("missing.example.com", None),
1661            FakeClock::new(),
1662        );
1663
1664        resolver
1665            .resolve(&query_for("missing.example.com"))
1666            .await
1667            .expect("nxdomain without soa still Ok");
1668        resolver
1669            .resolve(&query_for("missing.example.com"))
1670            .await
1671            .expect("served from cache using the fixed floor ttl");
1672        assert_eq!(
1673            calls.load(Ordering::SeqCst),
1674            1,
1675            "negative answer cached via floor"
1676        );
1677    }
1678
1679    #[tokio::test]
1680    async fn nodata_answer_is_cached_as_negative() {
1681        let policy = SplitDnsPolicy::builder()
1682            .default_group(UpstreamGroupId::new("g"))
1683            .build();
1684        let (resolver, calls) = resolver_with_counting_backend(
1685            policy,
1686            "g",
1687            nodata_answer("empty.example.com"),
1688            FakeClock::new(),
1689        );
1690
1691        resolver
1692            .resolve(&query_for("empty.example.com"))
1693            .await
1694            .expect("nodata is Ok(Message)");
1695        resolver
1696            .resolve(&query_for("empty.example.com"))
1697            .await
1698            .expect("served from cache");
1699        assert_eq!(calls.load(Ordering::SeqCst), 1, "nodata answer cached");
1700    }
1701
1702    #[tokio::test]
1703    async fn expired_cache_entry_triggers_a_fresh_backend_call() {
1704        let policy = SplitDnsPolicy::builder()
1705            .default_group(UpstreamGroupId::new("g"))
1706            .build();
1707        let clock = FakeClock::new();
1708        let (resolver, calls) =
1709            resolver_with_counting_backend(policy, "g", a_answer("example.com", 10), clock.clone());
1710
1711        resolver
1712            .resolve(&query_for("example.com"))
1713            .await
1714            .expect("first resolve populates cache");
1715        assert_eq!(calls.load(Ordering::SeqCst), 1);
1716
1717        clock.advance(Duration::from_secs(11));
1718
1719        resolver
1720            .resolve(&query_for("example.com"))
1721            .await
1722            .expect("expired entry re-queries the backend");
1723        assert_eq!(
1724            calls.load(Ordering::SeqCst),
1725            2,
1726            "ttl-expired entry is not served from cache"
1727        );
1728    }
1729
1730    #[tokio::test]
1731    async fn hook_use_overrides_the_static_group() {
1732        let hook_calls = Arc::new(AtomicUsize::new(0));
1733        let static_calls = Arc::new(AtomicUsize::new(0));
1734        let selected_calls = Arc::new(AtomicUsize::new(0));
1735        let resolver = Resolver::builder(
1736            SplitDnsPolicy::builder()
1737                .default_group(UpstreamGroupId::new("static"))
1738                .build(),
1739        )
1740        .backend(
1741            UpstreamGroupId::new("static"),
1742            CountingBackend {
1743                answer: answer_tagged(1),
1744                calls: static_calls.clone(),
1745            },
1746        )
1747        .backend(
1748            UpstreamGroupId::new("selected"),
1749            CountingBackend {
1750                answer: answer_tagged(2),
1751                calls: selected_calls.clone(),
1752            },
1753        )
1754        .route_hook(FixedHook {
1755            decision: Ok(RouteDecision::Use(UpstreamGroupId::new("selected"))),
1756            calls: hook_calls.clone(),
1757        })
1758        .build();
1759
1760        let answer = resolver.resolve(&query_for("example.com")).await.unwrap();
1761        assert_eq!(answer.header.id, 2);
1762        assert_eq!(hook_calls.load(Ordering::SeqCst), 1);
1763        assert_eq!(static_calls.load(Ordering::SeqCst), 0);
1764        assert_eq!(selected_calls.load(Ordering::SeqCst), 1);
1765    }
1766
1767    #[tokio::test]
1768    async fn hook_abstain_uses_the_static_group() {
1769        let backend_calls = Arc::new(AtomicUsize::new(0));
1770        let resolver = Resolver::builder(
1771            SplitDnsPolicy::builder()
1772                .default_group(UpstreamGroupId::new("static"))
1773                .build(),
1774        )
1775        .backend(
1776            UpstreamGroupId::new("static"),
1777            CountingBackend {
1778                answer: answer_tagged(3),
1779                calls: backend_calls.clone(),
1780            },
1781        )
1782        .route_hook(FixedHook {
1783            decision: Ok(RouteDecision::Abstain),
1784            calls: Arc::new(AtomicUsize::new(0)),
1785        })
1786        .build();
1787
1788        assert_eq!(
1789            resolver
1790                .resolve(&query_for("example.com"))
1791                .await
1792                .unwrap()
1793                .header
1794                .id,
1795            3
1796        );
1797        assert_eq!(backend_calls.load(Ordering::SeqCst), 1);
1798    }
1799
1800    #[tokio::test]
1801    async fn hook_observes_static_candidate_and_can_supply_a_route_without_one() {
1802        let static_groups = Arc::new(Mutex::new(Vec::new()));
1803        let static_resolver = Resolver::builder(
1804            SplitDnsPolicy::builder()
1805                .default_group(UpstreamGroupId::new("static"))
1806                .build(),
1807        )
1808        .backend(
1809            UpstreamGroupId::new("static"),
1810            fixed_backend(answer_tagged(30)),
1811        )
1812        .route_hook(RecordingHook {
1813            decision: RouteDecision::Abstain,
1814            static_groups: static_groups.clone(),
1815        })
1816        .build();
1817        assert_eq!(
1818            static_resolver
1819                .resolve(&query_for("static.example"))
1820                .await
1821                .unwrap()
1822                .header
1823                .id,
1824            30
1825        );
1826
1827        let dynamic_resolver = Resolver::builder(SplitDnsPolicy::builder().build())
1828            .backend(
1829                UpstreamGroupId::new("dynamic"),
1830                fixed_backend(answer_tagged(31)),
1831            )
1832            .route_hook(RecordingHook {
1833                decision: RouteDecision::Use(UpstreamGroupId::new("dynamic")),
1834                static_groups: static_groups.clone(),
1835            })
1836            .build();
1837        assert_eq!(
1838            dynamic_resolver
1839                .resolve(&query_for("dynamic.example"))
1840                .await
1841                .unwrap()
1842                .header
1843                .id,
1844            31
1845        );
1846        assert_eq!(
1847            *static_groups.lock().unwrap(),
1848            vec![Some(UpstreamGroupId::new("static")), None]
1849        );
1850    }
1851
1852    #[tokio::test]
1853    async fn hook_abstain_without_static_route_returns_no_route() {
1854        let backend_calls = Arc::new(AtomicUsize::new(0));
1855        let resolver = Resolver::builder(SplitDnsPolicy::builder().build())
1856            .backend(
1857                UpstreamGroupId::new("unused"),
1858                CountingBackend {
1859                    answer: answer_tagged(4),
1860                    calls: backend_calls.clone(),
1861                },
1862            )
1863            .route_hook(FixedHook {
1864                decision: Ok(RouteDecision::Abstain),
1865                calls: Arc::new(AtomicUsize::new(0)),
1866            })
1867            .build();
1868
1869        assert_eq!(
1870            resolver.resolve(&query_for("example.com")).await,
1871            Err(Error::NoRoute)
1872        );
1873        assert_eq!(backend_calls.load(Ordering::SeqCst), 0);
1874    }
1875
1876    #[tokio::test]
1877    async fn hook_selected_unknown_or_empty_group_returns_no_route_without_fallback() {
1878        for group in ["unknown", "empty"] {
1879            let static_calls = Arc::new(AtomicUsize::new(0));
1880            let builder = Resolver::builder(
1881                SplitDnsPolicy::builder()
1882                    .default_group(UpstreamGroupId::new("static"))
1883                    .build(),
1884            )
1885            .backend(
1886                UpstreamGroupId::new("static"),
1887                CountingBackend {
1888                    answer: answer_tagged(5),
1889                    calls: static_calls.clone(),
1890                },
1891            );
1892            let mut resolver = builder
1893                .route_hook(FixedHook {
1894                    decision: Ok(RouteDecision::Use(UpstreamGroupId::new(group))),
1895                    calls: Arc::new(AtomicUsize::new(0)),
1896                })
1897                .build();
1898            if group == "empty" {
1899                resolver
1900                    .backends
1901                    .insert(UpstreamGroupId::new("empty"), Vec::new());
1902            }
1903
1904            assert_eq!(
1905                resolver.resolve(&query_for("example.com")).await,
1906                Err(Error::NoRoute)
1907            );
1908            assert_eq!(
1909                static_calls.load(Ordering::SeqCst),
1910                0,
1911                "static backend must not receive a hook-selected {group} route"
1912            );
1913        }
1914    }
1915
1916    #[tokio::test]
1917    async fn hook_error_is_not_cached_retried_or_fallen_back() {
1918        let hook_calls = Arc::new(AtomicUsize::new(0));
1919        let backend_calls = Arc::new(AtomicUsize::new(0));
1920        let resolver = Resolver::builder(
1921            SplitDnsPolicy::builder()
1922                .default_group(UpstreamGroupId::new("static"))
1923                .build(),
1924        )
1925        .backend(
1926            UpstreamGroupId::new("static"),
1927            CountingBackend {
1928                answer: answer_tagged(6),
1929                calls: backend_calls.clone(),
1930            },
1931        )
1932        .route_hook(FixedHook {
1933            decision: Err(RouteHookError::new("policy unavailable")),
1934            calls: hook_calls.clone(),
1935        })
1936        .build();
1937
1938        for _ in 0..2 {
1939            assert_eq!(
1940                resolver.resolve(&query_for("example.com")).await,
1941                Err(Error::Hook("policy unavailable".to_string()))
1942            );
1943        }
1944        assert_eq!(hook_calls.load(Ordering::SeqCst), 2);
1945        assert_eq!(backend_calls.load(Ordering::SeqCst), 0);
1946        assert!(resolver.cache.lock().unwrap().is_empty());
1947    }
1948
1949    #[tokio::test]
1950    async fn cache_is_scoped_to_the_effective_hook_selected_group() {
1951        let first_calls = Arc::new(AtomicUsize::new(0));
1952        let second_calls = Arc::new(AtomicUsize::new(0));
1953        let resolver = Resolver::builder(SplitDnsPolicy::builder().build())
1954            .backend(
1955                UpstreamGroupId::new("first"),
1956                CountingBackend {
1957                    answer: a_answer("example.com", 300),
1958                    calls: first_calls.clone(),
1959                },
1960            )
1961            .backend(
1962                UpstreamGroupId::new("second"),
1963                CountingBackend {
1964                    answer: answer_tagged(8),
1965                    calls: second_calls.clone(),
1966                },
1967            )
1968            .route_hook(SequencedHook {
1969                decisions: Mutex::new(vec![
1970                    RouteDecision::Use(UpstreamGroupId::new("first")),
1971                    RouteDecision::Use(UpstreamGroupId::new("second")),
1972                    RouteDecision::Use(UpstreamGroupId::new("first")),
1973                ]),
1974            })
1975            .build();
1976
1977        let first = resolver
1978            .resolve(&query_for_type("example.com", RecordType::A, Class::In, 41))
1979            .await
1980            .unwrap();
1981        let second = resolver
1982            .resolve(&query_for_type("example.com", RecordType::A, Class::In, 42))
1983            .await
1984            .unwrap();
1985        let cached_first = resolver
1986            .resolve(&query_for_type("example.com", RecordType::A, Class::In, 43))
1987            .await
1988            .unwrap();
1989
1990        assert_eq!(first.answers[0].ttl, 300);
1991        assert_eq!(
1992            second.header.id, 8,
1993            "second route cannot reuse first route cache"
1994        );
1995        assert_eq!(cached_first.header.id, 43);
1996        assert_eq!(cached_first.questions, query_for("example.com").questions);
1997        assert_eq!(
1998            cached_first.answers, first.answers,
1999            "first route has its own cache hit"
2000        );
2001        assert_eq!(first_calls.load(Ordering::SeqCst), 1);
2002        assert_eq!(second_calls.load(Ordering::SeqCst), 1);
2003    }
2004
2005    #[tokio::test]
2006    async fn dropping_resolve_drops_the_hook_future_without_holding_cache_lock() {
2007        let entered = Arc::new(Notify::new());
2008        let dropped = Arc::new(AtomicBool::new(false));
2009        let resolver = Arc::new(
2010            Resolver::builder(SplitDnsPolicy::builder().build())
2011                .route_hook(PendingHook {
2012                    entered: entered.clone(),
2013                    dropped: dropped.clone(),
2014                })
2015                .build(),
2016        );
2017        let entered_wait = entered.notified();
2018        let task_resolver = resolver.clone();
2019        let task =
2020            tokio::spawn(async move { task_resolver.resolve(&query_for("example.com")).await });
2021
2022        entered_wait.await;
2023        assert!(
2024            resolver.cache.try_lock().is_ok(),
2025            "the resolver cache mutex is not held across hook await"
2026        );
2027        task.abort();
2028        assert!(task.await.unwrap_err().is_cancelled());
2029        assert!(
2030            dropped.load(Ordering::SeqCst),
2031            "hook future was dropped on cancellation"
2032        );
2033    }
2034
2035    fn query_for_type(name: &str, qtype: RecordType, qclass: Class, id: u16) -> Message {
2036        let mut query = query_for(name);
2037        query.header.id = id;
2038        query.questions[0].qtype = qtype;
2039        query.questions[0].qclass = qclass;
2040        query
2041    }
2042
2043    fn fake_ip_policy(name: &str) -> FakeIpPolicy {
2044        FakeIpPolicy::builder()
2045            .rule(DomainPattern::suffix(n(name)))
2046            .build()
2047    }
2048
2049    fn fake_ip_pool(clock: PoolClock) -> Arc<FakeIpPool> {
2050        Arc::new(
2051            FakeIpPool::builder()
2052                .ipv4_range(Ipv4Addr::new(198, 18, 0, 1), Ipv4Addr::new(198, 18, 0, 2))
2053                .ttl(Duration::from_secs(30))
2054                .clock(clock)
2055                .build()
2056                .unwrap(),
2057        )
2058    }
2059
2060    fn fake_ip_pool_ipv6(clock: PoolClock) -> Arc<FakeIpPool> {
2061        Arc::new(
2062            FakeIpPool::builder()
2063                .ipv6_range(
2064                    "2001:db8::1".parse().unwrap(),
2065                    "2001:db8::2".parse().unwrap(),
2066                )
2067                .ttl(Duration::from_secs(30))
2068                .clock(clock)
2069                .build()
2070                .unwrap(),
2071        )
2072    }
2073
2074    #[tokio::test]
2075    async fn fake_ip_a_answer_is_local_and_bypasses_upstream_and_cache() {
2076        let calls = Arc::new(AtomicUsize::new(0));
2077        let hook_calls = Arc::new(AtomicUsize::new(0));
2078        let backend = CountingBackend {
2079            answer: a_answer("example.test", 300),
2080            calls: calls.clone(),
2081        };
2082        let pool = fake_ip_pool(PoolClock::new());
2083        let resolver = Resolver::builder(
2084            SplitDnsPolicy::builder()
2085                .default_group(UpstreamGroupId::new("g"))
2086                .build(),
2087        )
2088        .backend(UpstreamGroupId::new("g"), backend)
2089        .fake_ip(pool, fake_ip_policy("example.test"))
2090        .route_hook(FixedHook {
2091            decision: Ok(RouteDecision::Use(UpstreamGroupId::new("g"))),
2092            calls: hook_calls.clone(),
2093        })
2094        .build();
2095
2096        let first = resolver
2097            .resolve(&query_for_type(
2098                "www.example.test",
2099                RecordType::A,
2100                Class::In,
2101                41,
2102            ))
2103            .await
2104            .unwrap();
2105        let second = resolver
2106            .resolve(&query_for_type(
2107                "www.example.test",
2108                RecordType::A,
2109                Class::In,
2110                42,
2111            ))
2112            .await
2113            .unwrap();
2114
2115        assert_eq!(calls.load(Ordering::SeqCst), 0);
2116        assert_eq!(
2117            hook_calls.load(Ordering::SeqCst),
2118            0,
2119            "Fake IP is terminal before hooks"
2120        );
2121        assert_eq!(first.header.id, 41);
2122        assert_eq!(second.header.id, 42, "synthetic answers are not cached");
2123        assert!(first.header.qr);
2124        assert_eq!(first.questions, query_for("www.example.test").questions);
2125        assert_eq!(first.answers[0].ttl, 30);
2126        assert_eq!(first.answers[0].rdata, second.answers[0].rdata);
2127    }
2128
2129    #[tokio::test]
2130    async fn fake_ip_disabled_family_returns_local_nodata() {
2131        let calls = Arc::new(AtomicUsize::new(0));
2132        let resolver = Resolver::builder(
2133            SplitDnsPolicy::builder()
2134                .default_group(UpstreamGroupId::new("g"))
2135                .build(),
2136        )
2137        .backend(
2138            UpstreamGroupId::new("g"),
2139            CountingBackend {
2140                answer: a_answer("example.test", 300),
2141                calls: calls.clone(),
2142            },
2143        )
2144        .fake_ip(
2145            fake_ip_pool(PoolClock::new()),
2146            fake_ip_policy("example.test"),
2147        )
2148        .build();
2149
2150        let answer = resolver
2151            .resolve(&query_for_type(
2152                "www.example.test",
2153                RecordType::Aaaa,
2154                Class::In,
2155                9,
2156            ))
2157            .await
2158            .unwrap();
2159
2160        assert_eq!(answer.header.rcode, Rcode::NoError);
2161        assert!(answer.answers.is_empty());
2162        assert_eq!(calls.load(Ordering::SeqCst), 0);
2163    }
2164
2165    #[tokio::test]
2166    async fn fake_ip_ptr_is_local_and_expires_with_its_mapping() {
2167        let pool_clock = PoolClock::new();
2168        let pool = fake_ip_pool(pool_clock.clone());
2169        let resolver = Resolver::builder(SplitDnsPolicy::builder().build())
2170            .fake_ip(pool.clone(), fake_ip_policy("example.test"))
2171            .build();
2172        let address = pool.allocate_ipv4(n("www.example.test")).unwrap();
2173        let reverse = format!(
2174            "{}.{}.{}.{}.in-addr.arpa",
2175            address.octets()[3],
2176            address.octets()[2],
2177            address.octets()[1],
2178            address.octets()[0]
2179        );
2180
2181        let found = resolver
2182            .resolve(&query_for_type(&reverse, RecordType::Ptr, Class::In, 11))
2183            .await
2184            .unwrap();
2185        assert_eq!(found.header.rcode, Rcode::NoError);
2186        assert_eq!(found.answers[0].rdata, RData::Ptr(n("www.example.test")));
2187        assert_eq!(found.answers[0].ttl, 30);
2188
2189        pool_clock.advance(Duration::from_secs(30));
2190        let expired = resolver
2191            .resolve(&query_for_type(&reverse, RecordType::Ptr, Class::In, 12))
2192            .await
2193            .unwrap();
2194        assert_eq!(expired.header.rcode, Rcode::NxDomain);
2195        assert!(expired.answers.is_empty());
2196    }
2197
2198    #[tokio::test]
2199    async fn fake_ip_answer_ttl_never_outlives_existing_mapping() {
2200        let pool_clock = PoolClock::new();
2201        let pool = fake_ip_pool(pool_clock.clone());
2202        let resolver = Resolver::builder(SplitDnsPolicy::builder().build())
2203            .fake_ip(pool.clone(), fake_ip_policy("example.test"))
2204            .build();
2205        pool.allocate_ipv4(n("www.example.test")).unwrap();
2206
2207        pool_clock.advance(Duration::from_secs(29));
2208        let answer = resolver
2209            .resolve(&query_for_type(
2210                "www.example.test",
2211                RecordType::A,
2212                Class::In,
2213                20,
2214            ))
2215            .await
2216            .unwrap();
2217
2218        assert_eq!(answer.answers[0].ttl, 1);
2219    }
2220
2221    #[tokio::test]
2222    async fn fake_ip_ptr_ttl_never_outlives_existing_mapping() {
2223        let pool_clock = PoolClock::new();
2224        let pool = fake_ip_pool(pool_clock.clone());
2225        let resolver = Resolver::builder(SplitDnsPolicy::builder().build())
2226            .fake_ip(pool.clone(), fake_ip_policy("example.test"))
2227            .build();
2228        let address = pool.allocate_ipv4(n("www.example.test")).unwrap();
2229        let reverse = format!(
2230            "{}.{}.{}.{}.in-addr.arpa",
2231            address.octets()[3],
2232            address.octets()[2],
2233            address.octets()[1],
2234            address.octets()[0]
2235        );
2236
2237        pool_clock.advance(Duration::from_secs(29));
2238        let answer = resolver
2239            .resolve(&query_for_type(&reverse, RecordType::Ptr, Class::In, 21))
2240            .await
2241            .unwrap();
2242
2243        assert_eq!(answer.answers[0].ttl, 1);
2244    }
2245
2246    #[tokio::test]
2247    async fn fake_ip_ipv6_ptr_is_local() {
2248        let pool = fake_ip_pool_ipv6(PoolClock::new());
2249        let resolver = Resolver::builder(SplitDnsPolicy::builder().build())
2250            .fake_ip(pool.clone(), fake_ip_policy("example.test"))
2251            .build();
2252        let address = pool.allocate_ipv6(n("www.example.test")).unwrap();
2253        let reverse = address
2254            .octets()
2255            .iter()
2256            .rev()
2257            .flat_map(|byte| [format!("{:x}", byte & 0x0f), format!("{:x}", byte >> 4)])
2258            .collect::<Vec<_>>()
2259            .join(".");
2260
2261        let answer = resolver
2262            .resolve(&query_for_type(
2263                &format!("{reverse}.ip6.arpa"),
2264                RecordType::Ptr,
2265                Class::In,
2266                22,
2267            ))
2268            .await
2269            .unwrap();
2270
2271        assert_eq!(answer.header.rcode, Rcode::NoError);
2272        assert_eq!(answer.answers[0].rdata, RData::Ptr(n("www.example.test")));
2273    }
2274
2275    #[tokio::test]
2276    async fn normal_queries_and_outside_reverse_ranges_still_use_upstream() {
2277        let calls = Arc::new(AtomicUsize::new(0));
2278        let pool = Arc::new(
2279            FakeIpPool::builder()
2280                .ipv4_range(Ipv4Addr::new(198, 18, 0, 1), Ipv4Addr::new(198, 18, 0, 2))
2281                .ttl(Duration::from_secs(u64::MAX))
2282                .clock(PoolClock::new())
2283                .build()
2284                .unwrap(),
2285        );
2286        let resolver = Resolver::builder(
2287            SplitDnsPolicy::builder()
2288                .default_group(UpstreamGroupId::new("g"))
2289                .build(),
2290        )
2291        .backend(
2292            UpstreamGroupId::new("g"),
2293            CountingBackend {
2294                answer: answer_tagged(77),
2295                calls: calls.clone(),
2296            },
2297        )
2298        .fake_ip(pool, fake_ip_policy("selected.test"))
2299        .build();
2300
2301        for query in [
2302            query_for_type("miss.test", RecordType::A, Class::In, 1),
2303            query_for_type("selected.test", RecordType::A, Class::Ch, 2),
2304            query_for_type("selected.test", RecordType::Txt, Class::In, 3),
2305            query_for_type("1.0.0.203.in-addr.arpa", RecordType::Ptr, Class::In, 4),
2306        ] {
2307            let answer = resolver.resolve(&query).await.unwrap();
2308            assert_eq!(answer.header.id, 77);
2309        }
2310        assert_eq!(calls.load(Ordering::SeqCst), 4);
2311    }
2312
2313    #[tokio::test]
2314    async fn unrepresentable_fake_ip_ttl_fails_before_allocation() {
2315        let pool = Arc::new(
2316            FakeIpPool::builder()
2317                .ipv4_range(Ipv4Addr::new(198, 18, 0, 1), Ipv4Addr::new(198, 18, 0, 2))
2318                .ttl(Duration::from_secs(u64::from(u32::MAX) + 1))
2319                .clock(PoolClock::new())
2320                .build()
2321                .unwrap(),
2322        );
2323        let resolver = Resolver::builder(SplitDnsPolicy::builder().build())
2324            .fake_ip(pool.clone(), fake_ip_policy("example.test"))
2325            .build();
2326
2327        assert_eq!(
2328            resolver
2329                .resolve(&query_for_type(
2330                    "www.example.test",
2331                    RecordType::A,
2332                    Class::In,
2333                    23
2334                ))
2335                .await,
2336            Err(Error::FakeIpTtlOutOfRange)
2337        );
2338        assert!(pool.snapshot().mappings.is_empty());
2339    }
2340
2341    #[tokio::test]
2342    async fn disabled_fake_ip_families_return_nodata_even_with_unrepresentable_ttl() {
2343        let calls = Arc::new(AtomicUsize::new(0));
2344        let pool = Arc::new(
2345            FakeIpPool::builder()
2346                .ipv4_range(Ipv4Addr::new(198, 18, 0, 1), Ipv4Addr::new(198, 18, 0, 2))
2347                .ttl(Duration::from_secs(u64::from(u32::MAX) + 1))
2348                .clock(PoolClock::new())
2349                .build()
2350                .unwrap(),
2351        );
2352        let resolver = Resolver::builder(
2353            SplitDnsPolicy::builder()
2354                .default_group(UpstreamGroupId::new("g"))
2355                .build(),
2356        )
2357        .backend(
2358            UpstreamGroupId::new("g"),
2359            CountingBackend {
2360                answer: answer_tagged(78),
2361                calls: calls.clone(),
2362            },
2363        )
2364        .fake_ip(pool.clone(), fake_ip_policy("example.test"))
2365        .build();
2366
2367        let aaaa = resolver
2368            .resolve(&query_for_type(
2369                "www.example.test",
2370                RecordType::Aaaa,
2371                Class::In,
2372                24,
2373            ))
2374            .await
2375            .unwrap();
2376        assert_eq!(aaaa.header.rcode, Rcode::NoError);
2377        assert!(aaaa.answers.is_empty());
2378        assert!(pool.snapshot().mappings.is_empty());
2379        assert_eq!(calls.load(Ordering::SeqCst), 0);
2380
2381        let ipv6_only_pool = Arc::new(
2382            FakeIpPool::builder()
2383                .ipv6_range(
2384                    "2001:db8::1".parse().unwrap(),
2385                    "2001:db8::2".parse().unwrap(),
2386                )
2387                .ttl(Duration::from_secs(u64::from(u32::MAX) + 1))
2388                .clock(PoolClock::new())
2389                .build()
2390                .unwrap(),
2391        );
2392        let ipv6_only_resolver = Resolver::builder(
2393            SplitDnsPolicy::builder()
2394                .default_group(UpstreamGroupId::new("g"))
2395                .build(),
2396        )
2397        .backend(
2398            UpstreamGroupId::new("g"),
2399            CountingBackend {
2400                answer: answer_tagged(79),
2401                calls: calls.clone(),
2402            },
2403        )
2404        .fake_ip(ipv6_only_pool.clone(), fake_ip_policy("example.test"))
2405        .build();
2406
2407        let a = ipv6_only_resolver
2408            .resolve(&query_for_type(
2409                "www.example.test",
2410                RecordType::A,
2411                Class::In,
2412                25,
2413            ))
2414            .await
2415            .unwrap();
2416        assert_eq!(a.header.rcode, Rcode::NoError);
2417        assert!(a.answers.is_empty());
2418        assert!(ipv6_only_pool.snapshot().mappings.is_empty());
2419        assert_eq!(calls.load(Ordering::SeqCst), 0);
2420    }
2421
2422    #[test]
2423    fn parses_canonical_ipv4_and_ipv6_reverse_names() {
2424        assert_eq!(
2425            parse_reverse_name(&n("4.3.2.1.in-addr.arpa")),
2426            Some(std::net::IpAddr::V4(Ipv4Addr::new(1, 2, 3, 4)))
2427        );
2428        assert_eq!(
2429            parse_reverse_name(&n("4.3.2.01.in-addr.arpa")),
2430            None,
2431            "non-canonical decimal labels are routed normally"
2432        );
2433        let reverse = format!("1.{}ip6.arpa", "0.".repeat(31));
2434        assert_eq!(
2435            parse_reverse_name(&n(&reverse)),
2436            Some(std::net::IpAddr::V6("::1".parse().unwrap()))
2437        );
2438    }
2439}