package traverse import ( "context" "fmt" "net" "sync" "github.com/hits/ExploreDNS/internal/dns" miekgdns "github.com/miekg/dns" ) type TraverserConfig struct { MaxDepth int QueryType uint16 RootConfig *dns.RootDiscoveryConfig QueryConfig *dns.QueryConfig RootAddrs []net.IP Hooks *TraverserHooks } func DefaultTraverserConfig() *TraverserConfig { return &TraverserConfig{ MaxDepth: DefaultMaxDepth, QueryType: dns.TypeA, RootConfig: nil, QueryConfig: nil, RootAddrs: nil, } } type TraversalResult struct { Referral *Referral Response *Response } 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 } cache := rootCache if ref.Parent != nil { cache = rootCache.Child() } 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 { 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, } } 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, WriteTimeout: 5, } if deadline, ok := ctx.Deadline(); ok { c.ReadTimeout = deadline.Sub(deadline) c.WriteTimeout = deadline.Sub(deadline) } 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 }