Files
ExploreDNS/internal/traverse/traverser.go
T
9aa85d8e5d
CI / test (pull_request) Failing after 2m13s
docs: comprehensive documentation for ExploreDNS
- Rewrite README.md with overview, features, installation (go install +
  build from source), quick start, full CLI flag reference table,
  output section descriptions, project structure, and development guide

- Add package-level doc comments to all five internal packages:
  config, dns, fingerprint, output, traverse (via doc.go or existing
  package-declaration files)

- Add GoDoc comments on every exported type, constant, function, and
  method across all packages:
  - internal/config: Config struct fields, all Parse*/Default/Validate
  - internal/dns: QueryConfig, Resolver, BasicResolver, CachingResolver,
    ExchangeFunc, RootServer, RootDiscoveryConfig, DecodedResponse,
    ResponseClassification, all exported helpers
  - internal/fingerprint: Fingerprinter, New, NewWithTimeout, Query,
    FingerprintAll
  - internal/traverse: Traverser, TraverserConfig, TraversalResult,
    Referral, ResolutionState, Response, ResponseType, InfoCache,
    Stack, TraverserHooks, EventStage, TraversalEvent, EventHandler,
    CircularReferralError, UnresolvableNameserverError
  - internal/output: Format, Config, Formatter, SummaryStats,
    NewFormatter, AttachHooks, RunTraversal, DefaultConfig,
    ComputeSummary

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Co-authored-by: multica-agent <github@multica.ai>
2026-06-08 04:01:51 +10:00

511 lines
12 KiB
Go

package traverse
import (
"context"
"fmt"
"net"
"sync"
"time"
"github.com/hits/ExploreDNS/internal/dns"
miekgdns "github.com/miekg/dns"
)
// TraverserConfig controls the behaviour of a Traverser.
type TraverserConfig struct {
// MaxDepth caps the traversal depth. Referrals at or beyond this depth
// are rejected and reported as errors.
MaxDepth int
// QueryType is the DNS record type requested at each step (e.g. dns.TypeA).
QueryType uint16
// RootConfig controls how root servers are discovered at startup.
RootConfig *dns.RootDiscoveryConfig
// QueryConfig controls UDP/TCP transport settings for each DNS query.
QueryConfig *dns.QueryConfig
// RootAddrs may be supplied directly to skip root discovery.
RootAddrs []net.IP
// Hooks receives events during traversal (progress, resolve, result).
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
}
// DefaultTraverserConfig returns a TraverserConfig with sensible defaults:
// max depth 20, query type A, fast mode on.
func DefaultTraverserConfig() *TraverserConfig {
return &TraverserConfig{
MaxDepth: DefaultMaxDepth,
QueryType: dns.TypeA,
RootConfig: nil,
QueryConfig: nil,
RootAddrs: nil,
Fast: true,
}
}
// TraversalResult pairs a Referral (the query that was attempted) with the
// Response (the outcome of that query). Response may be nil for referrals
// that were never processed (e.g. depth-limit rejections).
type TraversalResult struct {
Referral *Referral
Response *Response
}
// Traverser performs iterative DNS traversal from the root down to the target
// domain, following referrals and CNAME chains.
// Create one with NewTraverser; call Traverse to run a traversal.
type Traverser struct {
config *TraverserConfig
exchange dns.ExchangeFunc
visited map[string]bool
depth int
mu sync.Mutex
}
// NewTraverser creates a Traverser using cfg.
// When cfg is nil, DefaultTraverserConfig is used.
func NewTraverser(cfg *TraverserConfig) *Traverser {
if cfg == nil {
cfg = DefaultTraverserConfig()
}
return &Traverser{
config: cfg,
exchange: nil,
visited: make(map[string]bool),
depth: 0,
}
}
// SetExchange injects a custom exchange function, primarily for testing.
func (t *Traverser) SetExchange(fn dns.ExchangeFunc) {
t.exchange = fn
}
// SetHooks attaches traversal event hooks to the Traverser.
func (t *Traverser) SetHooks(hooks *TraverserHooks) {
if t.config == nil {
t.config = DefaultTraverserConfig()
}
t.config.Hooks = hooks
}
// Traverse performs an iterative DNS traversal for name, starting from the root
// servers. It returns all TraversalResults, including intermediate referrals
// and terminal outcomes. Hooks are called for each event during the traversal.
// The context can be used to cancel a long-running traversal.
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)
if err != nil {
return &Response{
Referral: ref,
Type: RespError,
}
}
if len(addrs) > 0 {
ref.Addresses = addrs
ref.State = StateResolved
} else {
return &Response{
Referral: ref,
Type: RespError,
}
}
}
}
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,
}
}
// ResolveNS resolves a nameserver hostname to its IP addresses by performing
// a fresh iterative traversal for that name, using cache to avoid repeated
// queries and visited to detect circular referrals.
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 {
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
}