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 }