From feb5c68b39a22ec0ace9d14ecc8298cd80ba3a77 Mon Sep 17 00:00:00 2001 From: Hansen IT Solutions Date: Sun, 7 Jun 2026 06:09:49 +0000 Subject: [PATCH] feat: implement referral resolution for nameserver names without glue records (HAN-381) (#7) Co-authored-by: Hansen IT Solutions Co-committed-by: Hansen IT Solutions --- internal/traverse/referral.go | 160 +++++++++++++++++++++++++++++ internal/traverse/referral_test.go | 48 +++++++++ internal/traverse/traverser.go | 155 +++++++++++++++++++++++++++- 3 files changed, 360 insertions(+), 3 deletions(-) diff --git a/internal/traverse/referral.go b/internal/traverse/referral.go index b6a553f..04be8d1 100644 --- a/internal/traverse/referral.go +++ b/internal/traverse/referral.go @@ -1,9 +1,12 @@ package traverse import ( + "context" + "fmt" "net" "strings" + "github.com/hits/ExploreDNS/internal/dns" miekgdns "github.com/miekg/dns" ) @@ -76,3 +79,160 @@ func (r *Referral) SetAddresses(addrs []net.IP) { r.State = StateUnresolved } } + +type CircularReferralError struct { + Name string + Chain []string +} + +func (e *CircularReferralError) Error() string { + return fmt.Sprintf("circular referral detected for %s: %v", e.Name, e.Chain) +} + +type UnresolvableNameserverError struct { + Name string + Reason string +} + +func (e *UnresolvableNameserverError) Error() string { + return fmt.Sprintf("unresolvable nameserver %s: %s", e.Name, e.Reason) +} + +func (r *Referral) Resolve(ctx context.Context, traverser *Traverser, cache *InfoCache, visited map[string]bool, depth int) error { + if r.HasAddresses() { + r.State = StateResolved + return nil + } + + if cache != nil { + if addrs := cache.LookupGlue(r.Name); len(addrs) > 0 { + r.Addresses = addrs + r.State = StateResolved + return nil + } + } + + if visited != nil { + if visited[r.Name] { + return &CircularReferralError{ + Name: r.Name, + Chain: getVisitedNames(visited), + } + } + visited[r.Name] = true + } + + if depth > DefaultMaxDepth { + return &UnresolvableNameserverError{ + Name: r.Name, + Reason: "max depth exceeded", + } + } + + roots, err := traverser.discoverRoots(ctx) + if err != nil { + return fmt.Errorf("root discovery: %w", err) + } + + initial := NewReferral(r.Name, dns.TypeA, ".", 0, 1.0, nil) + initial.Addresses = roots + initial.State = StateResolved + + stack := NewStack(DefaultMaxDepth) + stack.Push(initial) + + traversalCache := NewInfoCache(nil) + if visited != nil { + for name := range visited { + traversalCache.StoreGlue(name, []net.IP{}) + } + } + + var lastErr error + for { + select { + case <-ctx.Done(): + return fmt.Errorf("resolution cancelled: %w", ctx.Err()) + default: + } + + ref := stack.Pop() + if ref == nil { + break + } + + cacheForStep := traversalCache + if ref.Parent != nil { + cacheForStep = traversalCache.Child() + } + + resp := traverser.processReferral(ctx, ref, cacheForStep) + + if resp.Type == RespAnswer { + if len(resp.Decoded.Answers) > 0 { + var addrs []net.IP + 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 { + r.Addresses = addrs + r.State = StateResolved + if cache != nil { + cache.StoreGlue(r.Name, addrs) + } + return nil + } + } + } + + if resp.Type == RespNXDOMAIN { + lastErr = &UnresolvableNameserverError{ + Name: r.Name, + Reason: "NXDOMAIN", + } + break + } + + if resp.Type == RespSERVFAIL || resp.Type == RespError { + lastErr = fmt.Errorf("server error resolving %s: %s", r.Name, 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: r.Name, + Reason: "max depth exceeded during resolution", + } + } + } + } + } + + if lastErr != nil { + return lastErr + } + + return &UnresolvableNameserverError{ + Name: r.Name, + Reason: "resolution exhausted without answer", + } +} + +func getVisitedNames(visited map[string]bool) []string { + var names []string + for name := range visited { + names = append(names, name) + } + return names +} diff --git a/internal/traverse/referral_test.go b/internal/traverse/referral_test.go index 572d12d..0adfa61 100644 --- a/internal/traverse/referral_test.go +++ b/internal/traverse/referral_test.go @@ -109,3 +109,51 @@ func TestResolutionStateString(t *testing.T) { }) } } + +func TestCircularReferralError(t *testing.T) { + err := &CircularReferralError{ + Name: "ns.example.com.", + Chain: []string{"ns1.example.com.", "ns2.example.com."}, + } + + expected := "circular referral detected for ns.example.com.: [ns1.example.com. ns2.example.com.]" + if err.Error() != expected { + t.Errorf("Error() = %q, want %q", err.Error(), expected) + } +} + +func TestUnresolvableNameserverError(t *testing.T) { + err := &UnresolvableNameserverError{ + Name: "ns.example.com.", + Reason: "NXDOMAIN", + } + + expected := "unresolvable nameserver ns.example.com.: NXDOMAIN" + if err.Error() != expected { + t.Errorf("Error() = %q, want %q", err.Error(), expected) + } +} + +func TestGetVisitedNames(t *testing.T) { + visited := map[string]bool{ + "ns1.example.com.": true, + "ns2.example.com.": true, + "ns3.example.com.": true, + } + + names := getVisitedNames(visited) + if len(names) != 3 { + t.Errorf("got %d names, want 3", len(names)) + } + + seen := make(map[string]bool) + for _, name := range names { + if seen[name] { + t.Errorf("duplicate name: %s", name) + } + seen[name] = true + if !visited[name] { + t.Errorf("unexpected name: %s", name) + } + } +} diff --git a/internal/traverse/traverser.go b/internal/traverse/traverser.go index f98078d..b262f69 100644 --- a/internal/traverse/traverser.go +++ b/internal/traverse/traverser.go @@ -36,6 +36,9 @@ type TraversalResult struct { type Traverser struct { config *TraverserConfig exchange dns.ExchangeFunc + visited map[string]bool + depth int + mu sync.Mutex } func NewTraverser(cfg *TraverserConfig) *Traverser { @@ -45,6 +48,8 @@ func NewTraverser(cfg *TraverserConfig) *Traverser { return &Traverser{ config: cfg, exchange: nil, + visited: make(map[string]bool), + depth: 0, } } @@ -155,13 +160,32 @@ func (t *Traverser) discoverRoots(ctx context.Context) ([]net.IP, error) { 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 { - return &Response{ - Referral: ref, - Type: RespError, + 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, + } } } } @@ -179,6 +203,131 @@ func (t *Traverser) processReferral(ctx context.Context, ref *Referral, cache *I } } +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() + } + + resp := t.processReferral(ctx, current, cacheForStep) + + 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