package traverse import ( "context" "fmt" "net" "strings" "sync" "time" "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 before the traversal gives up. 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. RootAddrs []net.IP // Hooks provides optional callbacks for traversal events. Hooks *TraverserHooks // Fast controls cache sharing across branches. When true (default), child // branches inherit glue discovered by earlier branches via the shared root // cache, trading accuracy for speed. When false, each branch gets a // completely independent cache — slower but results are not contaminated by // sibling branch observations. Fast bool } func DefaultTraverserConfig() *TraverserConfig { return &TraverserConfig{ MaxDepth: DefaultMaxDepth, QueryType: dns.TypeA, RootConfig: nil, QueryConfig: nil, RootAddrs: nil, Fast: true, } } // TraversalResult pairs a Referral with the Response received when it was processed. type TraversalResult struct { Referral *Referral Response *Response } // Traverser performs an exhaustive iterative DNS traversal starting from the // root servers. Create one via NewTraverser and call Traverse to start a run. type Traverser struct { config *TraverserConfig exchange dns.ExchangeFunc visited map[string]bool depth int mu sync.Mutex } func NewTraverser(cfg *TraverserConfig) *Traverser { if cfg == nil { cfg = DefaultTraverserConfig() } return &Traverser{ config: cfg, exchange: nil, visited: make(map[string]bool), depth: 0, } } func (t *Traverser) SetExchange(fn dns.ExchangeFunc) { t.exchange = fn } func (t *Traverser) SetHooks(hooks *TraverserHooks) { if t.config == nil { t.config = DefaultTraverserConfig() } t.config.Hooks = hooks } func (t *Traverser) Traverse(ctx context.Context, name string) ([]TraversalResult, error) { name = miekgdns.Fqdn(name) roots, err := t.discoverRoots(ctx) if err != nil { return nil, fmt.Errorf("root discovery: %w", err) } initial := NewReferral(name, t.config.QueryType, ".", 0, 1.0, nil) initial.Addresses = roots stack := NewStack(t.config.MaxDepth) stack.Push(initial) rootCache := NewInfoCache(nil) var ( mu sync.Mutex results []TraversalResult ) for { select { case <-ctx.Done(): return results, fmt.Errorf("traversal cancelled: %w", ctx.Err()) default: } ref := stack.Pop() if ref == nil { break } var cache *InfoCache if t.config.Fast { // Fast mode: inherit glue from the shared root cache so earlier // branch discoveries are visible to later branches. cache = rootCache if ref.Parent != nil { cache = rootCache.Child() } } else { // Non-fast mode: every referral gets its own independent cache so // no cross-branch glue is reused, ensuring each path is resolved // from scratch. cache = NewInfoCache(nil) } if t.config.Hooks != nil { t.config.Hooks.emit(EventStart, TraversalResult{Referral: ref}, false) } resp := t.processReferral(ctx, ref, cache) result := TraversalResult{Referral: ref, Response: resp} if t.config.Hooks != nil { t.config.Hooks.emit(EventComplete, result, false) } mu.Lock() results = append(results, result) mu.Unlock() if resp.IsTerminal() { continue } if resp.Type == RespReferral { children := resp.ChildReferrals() for _, child := range children { if !stack.Push(child) { mu.Lock() results = append(results, TraversalResult{ Referral: child, Response: &Response{ Referral: child, Type: RespError, }, }) mu.Unlock() } } } if resp.Type == RespCNAMEFollow { follow := resp.CNAMEFollowReferral() if follow != nil { // Detect CNAME loop: target name already appears in the ancestor chain. if follow.Parent != nil && follow.Parent.IsNameInChain(follow.Name) { mu.Lock() results = append(results, TraversalResult{ Referral: follow, Response: &Response{ Referral: follow, Type: RespCNAMELoop, ErrorMessage: fmt.Sprintf("CNAME loop detected: %s already in traversal chain", follow.Name), }, }) mu.Unlock() } else if !stack.Push(follow) { mu.Lock() results = append(results, TraversalResult{ Referral: follow, Response: &Response{ Referral: follow, Type: RespError, }, }) mu.Unlock() } } } } return results, nil } func (t *Traverser) discoverRoots(ctx context.Context) ([]net.IP, error) { if len(t.config.RootAddrs) > 0 { return t.config.RootAddrs, nil } servers, err := dns.DiscoverRoots(ctx, t.config.RootConfig) if err != nil { return nil, err } var addrs []net.IP for _, srv := range servers { addrs = append(addrs, srv.AllIPs(false)...) } return addrs, nil } func (t *Traverser) processReferral(ctx context.Context, ref *Referral, cache *InfoCache) *Response { if !ref.HasAddresses() { t.mu.Lock() visitedCopy := make(map[string]bool) for k, v := range t.visited { visitedCopy[k] = v } t.mu.Unlock() ref.Addresses = t.resolveGlueViaSystem(ctx, ref.Name, cache) if len(ref.Addresses) > 0 { ref.State = StateResolved } else { addrs, err := t.ResolveNS(ctx, ref.Name, cache, visitedCopy, t.depth) nsName := strings.TrimSuffix(ref.Bailiwick, ".") if nsName == "" || nsName == "." { nsName = strings.TrimSuffix(ref.Name, ".") } if err != nil { return &Response{ Referral: ref, Type: RespNSResolutionFailed, ErrorMessage: fmt.Sprintf("nameserver %s could not be resolved", nsName), } } if len(addrs) > 0 { ref.Addresses = addrs ref.State = StateResolved } else { return &Response{ Referral: ref, Type: RespNSResolutionFailed, ErrorMessage: fmt.Sprintf("nameserver %s could not be resolved", nsName), } } } } for _, addr := range ref.Addresses { resp := t.queryServer(ctx, ref, addr, cache) if resp != nil && resp.Type != RespSERVFAIL { return resp } } return &Response{ Referral: ref, Type: RespSERVFAIL, } } func (t *Traverser) ResolveNS(ctx context.Context, nsName string, cache *InfoCache, visited map[string]bool, depth int) ([]net.IP, error) { if cache != nil { if addrs := cache.LookupGlue(nsName); len(addrs) > 0 { return addrs, nil } } if visited != nil { if visited[nsName] { return nil, &CircularReferralError{ Name: nsName, Chain: getVisitedNames(visited), } } visited[nsName] = true } if depth > DefaultMaxDepth { return nil, &UnresolvableNameserverError{ Name: nsName, Reason: "max depth exceeded", } } roots, err := t.discoverRoots(ctx) if err != nil { return nil, fmt.Errorf("root discovery: %w", err) } ref := NewReferral(nsName, dns.TypeA, ".", 0, 1.0, nil) ref.Addresses = roots ref.State = StateResolved traversalCache := NewInfoCache(nil) if visited != nil { for name := range visited { traversalCache.StoreGlue(name, []net.IP{}) } } var addrs []net.IP var lastErr error stack := NewStack(DefaultMaxDepth) stack.Push(ref) for { select { case <-ctx.Done(): return nil, fmt.Errorf("resolution cancelled: %w", ctx.Err()) default: } current := stack.Pop() if current == nil { break } cacheForStep := traversalCache if current.Parent != nil { cacheForStep = traversalCache.Child() } if t.config.Hooks != nil { t.config.Hooks.emit(EventStart, TraversalResult{Referral: current}, true) } resp := t.processReferral(ctx, current, cacheForStep) if t.config.Hooks != nil { t.config.Hooks.emit(EventComplete, TraversalResult{Referral: current, Response: resp}, true) } if resp.Type == RespAnswer && len(resp.Decoded.Answers) > 0 { for _, rr := range resp.Decoded.Answers { if a, ok := rr.(*miekgdns.A); ok { addrs = append(addrs, a.A) } if aaaa, ok := rr.(*miekgdns.AAAA); ok { addrs = append(addrs, aaaa.AAAA) } } if len(addrs) > 0 { if cache != nil { cache.StoreGlue(nsName, addrs) } return addrs, nil } } if resp.Type == RespNXDOMAIN { lastErr = &UnresolvableNameserverError{ Name: nsName, Reason: "NXDOMAIN", } break } if resp.Type == RespSERVFAIL || resp.Type == RespError || resp.Type == RespNSResolutionFailed { lastErr = fmt.Errorf("server error resolving %s: %s", nsName, resp.Type) continue } if resp.Type == RespReferral { children := resp.ChildReferrals() for _, child := range children { if visited != nil && visited[child.Name] { continue } if !stack.Push(child) { lastErr = &UnresolvableNameserverError{ Name: nsName, Reason: "max depth exceeded during resolution", } } } } } if len(addrs) > 0 { return addrs, nil } if lastErr != nil { return nil, lastErr } return nil, &UnresolvableNameserverError{ Name: nsName, Reason: "resolution exhausted without answer", } } func (t *Traverser) queryServer(ctx context.Context, ref *Referral, server net.IP, cache *InfoCache) *Response { var msg *miekgdns.Msg var err error if t.exchange != nil { msg, err = t.iterativeQueryWithExchange(ctx, server, ref.Name, ref.Qtype) } else { msg, err = dns.Query(ctx, server, ref.Name, ref.Qtype, t.config.QueryConfig) if err == nil { msg = t.ensureRDFalse(msg, server, ref.Name, ref.Qtype, t.config.QueryConfig) } } if err != nil { return &Response{ Referral: ref, Server: server, Type: RespError, } } resp := NewResponse(ref, server, cache) resp.Process(msg) return resp } func (t *Traverser) iterativeQueryWithExchange(ctx context.Context, server net.IP, name string, qtype uint16) (*miekgdns.Msg, error) { if t.config.QueryConfig == nil { return dns.IterativeQueryWithExchange(ctx, server, name, qtype, nil, t.exchange) } return dns.IterativeQueryWithExchange(ctx, server, name, qtype, t.config.QueryConfig, t.exchange) } func (t *Traverser) ensureRDFalse(msg *miekgdns.Msg, server net.IP, name string, qtype uint16, cfg *dns.QueryConfig) *miekgdns.Msg { if msg != nil && msg.RecursionDesired { if t.exchange != nil { ctx := context.Background() var err error msg, err = t.iterativeQueryWithExchange(ctx, server, name, qtype) if err != nil { return nil } return msg } msg.RecursionDesired = false } return msg } func (t *Traverser) resolveGlueViaSystem(ctx context.Context, name string, cache *InfoCache) []net.IP { if cache != nil { if addrs := cache.LookupGlue(name); len(addrs) > 0 { return addrs } } c := &miekgdns.Client{ Net: "udp", ReadTimeout: 5 * time.Second, WriteTimeout: 5 * time.Second, } if deadline, ok := ctx.Deadline(); ok { remaining := time.Until(deadline) if remaining <= 0 { return nil } c.ReadTimeout = remaining c.WriteTimeout = remaining } fqdn := miekgdns.Fqdn(name) aMsg, _, err := c.ExchangeContext(ctx, newAQuery(fqdn), "127.0.0.1:53") if err == nil { var addrs []net.IP for _, rr := range aMsg.Answer { if a, ok := rr.(*miekgdns.A); ok { addrs = append(addrs, a.A) } } if len(addrs) > 0 { if cache != nil { cache.StoreGlue(name, addrs) } return addrs } } return nil } func newAQuery(name string) *miekgdns.Msg { m := new(miekgdns.Msg) m.SetQuestion(name, miekgdns.TypeA) m.RecursionDesired = true return m }