Files
ExploreDNS/internal/traverse/traverser.go
Gary HansenandClaude Fable 5 d71c7fbef2 feat: rework engine and CLI for dnstraverse parity
Port the traversal engine to the Ruby dnstraverse model so behaviour and
output match dns.squish.net:

- dns: single RD=0 query path (RD=1 only for upstream root discovery),
  per-run packet cache, EDNS0 512-fallback with warnings, UDP->TCP on
  truncation; fix --retries 0 and --root-server IP-literal handling;
  drop all hardcoded 127.0.0.1:53 resolvers
- traverse: hierarchical per-branch InfoCache, 7-step response
  classification with the full 10-status vocabulary, bailiwick
  partitioning, strictly-deeper lame-referral rule, refid grammar with
  .0 resolve subtrees and childset digits, per-IP branching at 1/n
  weight, cache-based glue resolution with noglue/loop dead ends, CNAME
  restarts from the deepest cached zone, fast-mode memoization,
  probability aggregation with Ruby-identical stats keys (sums to 1.0)
- output: byte-for-byte reference text format pinned by a golden test,
  reference CLI defaults, working --quiet/--show-X=false, TTY-aware
  colour, deduplicated deterministic JSON
- web: adapt API/SPA to the new engine, SSE events carry refid/status,
  fix subscribe/snapshot duplicate-event race and a statusCls TDZ bug,
  align SPA type list with the backend
- delete the old engine and dead code (net -4,350 lines)

Verified against live runs of the reference Ruby engine across five
domains (answers, NXDOMAIN, null MX, CNAME restart, glueless resolve)
with no divergences beyond the documented typo fixes.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-07 21:42:06 +10:00

342 lines
9.3 KiB
Go

package traverse
import (
"context"
"fmt"
"net"
"sort"
"strings"
"gitea.hansenits.com.au/hits/ExploreDNS/internal/dns"
miekgdns "github.com/miekg/dns"
)
// TraverserConfig configures the behaviour of a Traverser.
type TraverserConfig struct {
// MaxDepth is the maximum referral depth (non-zero refid components)
// before a "Maxdepth N exceeded" exception is injected.
MaxDepth int
// QueryType is the DNS record type to query (e.g. dns.TypeA).
QueryType uint16
// RootConfig controls how root servers are discovered.
RootConfig *dns.RootDiscoveryConfig
// QueryConfig controls per-query transport parameters.
QueryConfig *dns.QueryConfig
// RootAddrs is an optional pre-seeded list of root server IP addresses.
// When non-empty, root discovery via RootConfig is skipped and each
// address becomes one root (named by its address).
RootAddrs []net.IP
// Hooks provides optional callbacks for traversal events.
Hooks *TraverserHooks
// Fast enables the completed-referral memo (traverser.rb @answered):
// a referral identical to an earlier completed one (same qname/qclass/
// qtype/server and per-IP weights) is replaced by it instead of being
// walked again. Non-fast mode re-walks every branch.
Fast bool
}
func DefaultTraverserConfig() *TraverserConfig {
return &TraverserConfig{
MaxDepth: DefaultMaxDepth,
QueryType: dns.TypeA,
Fast: true,
}
}
// Traverser drives the traversal: it owns the packet-cached query client,
// the fast-mode memo and the explicit stack loop (traverser.rb).
type Traverser struct {
config *TraverserConfig
client *dns.Client
exchange dns.ExchangeFunc
// answered is the fast-mode memo of completed referrals.
answered map[string]*Referral
// seen maps every server name encountered to its IP addresses.
seen map[string][]string
// roots memoises root discovery so Roots() and Run() share one lookup.
roots []StartServer
}
func NewTraverser(cfg *TraverserConfig) *Traverser {
if cfg == nil {
cfg = DefaultTraverserConfig()
}
if cfg.MaxDepth <= 0 {
cfg.MaxDepth = DefaultMaxDepth
}
if cfg.QueryType == 0 {
cfg.QueryType = dns.TypeA
}
return &Traverser{
config: cfg,
client: dns.NewClient(cfg.QueryConfig, nil),
answered: make(map[string]*Referral),
seen: make(map[string][]string),
}
}
// SetExchange injects a mock wire exchange into the single query path (both
// traversal queries and root discovery); tests use this so no packets leave
// the process.
func (t *Traverser) SetExchange(fn dns.ExchangeFunc) {
t.exchange = fn
t.client = dns.NewClient(t.config.QueryConfig, fn)
}
func (t *Traverser) SetHooks(hooks *TraverserHooks) {
t.config.Hooks = hooks
}
// ServersEncountered returns every server name seen during the run mapped to
// its known IP addresses (traverser.rb servers_encountered).
func (t *Traverser) ServersEncountered() map[string][]string {
return t.seen
}
// Roots performs (and memoises) root discovery, returning the start servers
// the traversal will begin from. Callers may use it before Run to report the
// initial root; Run reuses the memoised result.
func (t *Traverser) Roots(ctx context.Context) ([]StartServer, error) {
if t.roots == nil {
roots, err := t.rootStartServers(ctx)
if err != nil {
return nil, fmt.Errorf("root discovery: %w", err)
}
t.roots = roots
}
return t.roots, nil
}
// Run traverses the DNS for name and returns the synthetic rootroot node
// (never displayed) whose Stats aggregate every leaf outcome; the per-leaf
// probabilities sum to 1.0.
func (t *Traverser) Run(ctx context.Context, name string) (*Referral, error) {
roots, err := t.Roots(ctx)
if err != nil {
return nil, err
}
cache := NewInfoCache(nil)
cache.AddHints("", roots)
root := &Referral{
RefID: "",
Qname: canonicalName(toASCII(name)),
Qclass: miekgdns.ClassINET,
Qtype: t.config.QueryType,
NSAType: dns.TypeA,
Server: "",
Bailiwick: "",
InfoCache: cache,
Status: RefStatusNormal,
Responses: make(map[string]*ServerResponse),
Children: make(map[string][]*Referral),
ServerWeights: make(map[string]float64),
client: t.client,
maxdepth: t.config.MaxDepth,
}
t.config.Hooks.emit(StageNew, root, "")
if err := t.run(ctx, root); err != nil {
return root, err
}
return root, nil
}
// stack markers mirroring Ruby's :calc_resolve / :calc_answer placeholders:
// the referral is revisited after its resolves/children finished, giving
// post-order statistics calculation without recursion.
type stackMarker int
const (
markerNone stackMarker = iota
markerCalcResolve
markerCalcAnswer
)
type stackEntry struct {
ref *Referral
marker stackMarker
}
func (t *Traverser) run(ctx context.Context, root *Referral) error {
stack := []stackEntry{{ref: root}}
pop := func() stackEntry {
e := stack[len(stack)-1]
stack = stack[:len(stack)-1]
return e
}
for len(stack) > 0 {
select {
case <-ctx.Done():
return fmt.Errorf("traversal cancelled: %w", ctx.Err())
default:
}
e := pop()
r := e.ref
switch e.marker {
case markerCalcResolve:
r.resolveCalculate()
t.config.Hooks.emit(StageResolve, r, "")
stack = append(stack, stackEntry{ref: r}) // now needs processing
continue
case markerCalcAnswer:
r.answerCalculate()
t.config.Hooks.emit(StageAnswer, r, "")
if t.config.Fast && r.Status == RefStatusNormal && !hasLameResponse(r) {
t.answered[fastKey(r)] = r
}
if !r.IsRootRoot() {
t.recordSeen(r)
}
continue
}
// A new item. Fast mode: an identical completed referral replaces
// this one wholesale. noglue/loop nodes are excluded because their
// stats carry node-specific attributes and are cheap to recreate.
if t.config.Fast && r.Parent != nil {
if memo, ok := t.answered[fastKey(r)]; ok && !r.isNoGlue() && !r.isLoop() {
r.Parent.replaceChild(r, memo)
t.config.Hooks.emit(StageAnswerFast, r, memo.RefID)
continue
}
}
t.config.Hooks.emit(StageStart, r, "")
if !r.Resolved() {
// Push the resolve subtree with a calc_resolve placeholder so the
// weights are folded in once every resolve leaf completed.
stack = append(stack, stackEntry{ref: r, marker: markerCalcResolve})
resolves, err := r.resolve()
if err != nil {
return err
}
for _, c := range resolves {
t.config.Hooks.emit(StageNew, c, "")
}
for i := len(resolves) - 1; i >= 0; i-- {
stack = append(stack, stackEntry{ref: resolves[i]})
}
continue
}
stack = append(stack, stackEntry{ref: r, marker: markerCalcAnswer})
childrenSets, err := r.process(ctx)
if err != nil {
return err
}
seenParentIP := make(map[string]bool)
var flat []*Referral
for _, set := range childrenSets {
for _, c := range set {
if len(childrenSets) > 1 && !seenParentIP[c.ParentIP] {
t.config.Hooks.emit(StageNewReferralSet, c, "")
seenParentIP[c.ParentIP] = true
}
stage, earlier := StageNew, ""
if t.config.Fast {
if memo, ok := t.answered[fastKey(c)]; ok {
stage, earlier = StageNewFast, memo.RefID
}
}
t.config.Hooks.emit(stage, c, earlier)
flat = append(flat, c)
}
}
for i := len(flat) - 1; i >= 0; i-- {
stack = append(stack, stackEntry{ref: flat[i]})
}
}
return nil
}
// fastKey is the fast-mode memo key (traverser.rb): qname/qclass/qtype/
// server plus the per-IP weights, lowercased.
func fastKey(r *Referral) string {
return strings.ToLower(fmt.Sprintf("%s:%s:%s:%s:%s",
r.Qname, ClassToString(r.Qclass), TypeToString(r.Qtype), r.Server, r.TxtIPsVerbose()))
}
func hasLameResponse(r *Referral) bool {
for _, resp := range r.Responses {
if resp.Status == StatusReferralLame {
return true
}
}
return false
}
func (t *Traverser) recordSeen(r *Referral) {
name := strings.ToLower(r.Server)
existing := t.seen[name]
for _, ip := range r.IPsAsArray() {
found := false
for _, have := range existing {
if have == ip {
found = true
break
}
}
if !found {
existing = append(existing, ip)
}
}
t.seen[name] = existing
}
// rootStartServers returns the root servers as start-server hints: either
// the pre-seeded RootAddrs or the servers found via root discovery (one by
// default, all of them with AllRoots). IPv4 only, like the reference.
func (t *Traverser) rootStartServers(ctx context.Context) ([]StartServer, error) {
if len(t.config.RootAddrs) > 0 {
var out []StartServer
for _, ip := range t.config.RootAddrs {
if v4 := ip.To4(); v4 != nil {
out = append(out, StartServer{Name: v4.String(), IPs: []string{v4.String()}})
}
}
if len(out) == 0 {
return nil, fmt.Errorf("no usable IPv4 root addresses")
}
return out, nil
}
rootCfg := t.config.RootConfig
if t.exchange != nil {
var cp dns.RootDiscoveryConfig
if rootCfg != nil {
cp = *rootCfg
}
cp.Exchange = t.exchange
rootCfg = &cp
}
servers, err := dns.DiscoverRoots(ctx, rootCfg)
if err != nil {
return nil, err
}
var out []StartServer
for _, srv := range servers {
var ips []string
for _, ip := range srv.IPv4 {
ips = append(ips, ip.String())
}
if len(ips) == 0 {
continue
}
out = append(out, StartServer{Name: canonicalName(srv.Name), IPs: ips})
}
if len(out) == 0 {
return nil, fmt.Errorf("no root servers with IPv4 addresses")
}
sort.Slice(out, func(i, j int) bool { return out[i].Name < out[j].Name })
return out, nil
}