From fe1afe2a9789831dacb48c6f5d260a120a434f07 Mon Sep 17 00:00:00 2001 From: Multica Agent Date: Sun, 7 Jun 2026 17:24:00 +0000 Subject: [PATCH] Phase 4.1: Error handling, edge cases, and robustness (#11) --- cmd/exploredns/main.go | 1 + go.mod | 7 +- go.sum | 2 + internal/dns/decode.go | 46 ++ internal/dns/query.go | 18 +- internal/dns/robustness_test.go | 139 +++++ internal/output/stats.go | 6 + internal/output/text.go | 16 +- internal/traverse/referral.go | 42 +- internal/traverse/response.go | 45 +- internal/traverse/robustness_test.go | 768 +++++++++++++++++++++++++++ internal/traverse/traverser.go | 50 +- 12 files changed, 1118 insertions(+), 22 deletions(-) create mode 100644 internal/dns/robustness_test.go create mode 100644 internal/traverse/robustness_test.go diff --git a/cmd/exploredns/main.go b/cmd/exploredns/main.go index 6bcdf43..bd50c42 100644 --- a/cmd/exploredns/main.go +++ b/cmd/exploredns/main.go @@ -164,6 +164,7 @@ func main() { QueryType: queryTypeValue, RootConfig: rootConfig, QueryConfig: queryConfig, + Fast: cfg.Fast, } traverser := traverse.NewTraverser(traverserConfig) diff --git a/go.mod b/go.mod index 29ccbfa..04916b9 100644 --- a/go.mod +++ b/go.mod @@ -2,12 +2,15 @@ module github.com/hits/ExploreDNS go 1.25.6 -require github.com/miekg/dns v1.1.72 +require ( + github.com/miekg/dns v1.1.72 + golang.org/x/net v0.48.0 +) require ( golang.org/x/mod v0.31.0 // indirect - golang.org/x/net v0.48.0 // indirect golang.org/x/sync v0.19.0 // indirect golang.org/x/sys v0.39.0 // indirect + golang.org/x/text v0.32.0 // indirect golang.org/x/tools v0.40.0 // indirect ) diff --git a/go.sum b/go.sum index 92364b3..04b719e 100644 --- a/go.sum +++ b/go.sum @@ -10,5 +10,7 @@ golang.org/x/sync v0.19.0 h1:vV+1eWNmZ5geRlYjzm2adRgW2/mcpevXNg50YZtPCE4= golang.org/x/sync v0.19.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI= golang.org/x/sys v0.39.0 h1:CvCKL8MeisomCi6qNZ+wbb0DN9E5AATixKsvNtMoMFk= golang.org/x/sys v0.39.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks= +golang.org/x/text v0.32.0 h1:ZD01bjUt1FQ9WJ0ClOL5vxgxOI/sVCNgX1YtKwcY0mU= +golang.org/x/text v0.32.0/go.mod h1:o/rUWzghvpD5TXrTIBuJU77MTaN0ljMWE47kxGJQ7jY= golang.org/x/tools v0.40.0 h1:yLkxfA+Qnul4cs9QA3KnlFu0lVmd8JJfoq+E41uSutA= golang.org/x/tools v0.40.0/go.mod h1:Ik/tzLRlbscWpqqMRjyWYDisX8bG13FrdXp3o4Sr9lc= diff --git a/internal/dns/decode.go b/internal/dns/decode.go index a909f24..3e868f1 100644 --- a/internal/dns/decode.go +++ b/internal/dns/decode.go @@ -2,6 +2,7 @@ package dns import ( "fmt" + "strings" "github.com/miekg/dns" ) @@ -14,6 +15,8 @@ const ( ResponseNODATA ResponseNXDOMAIN ResponseSERVFAIL + ResponseREFUSED + ResponseNOTIMPL ResponseOther ) @@ -29,6 +32,10 @@ func (rc ResponseClassification) String() string { return "nxdomain" case ResponseSERVFAIL: return "servfail" + case ResponseREFUSED: + return "refused" + case ResponseNOTIMPL: + return "notimp" default: return "other" } @@ -45,6 +52,13 @@ type DecodedResponse struct { Authority []dns.RR Additional []dns.RR CNAMEChain []string + DNAMEMappings []DNAMEMapping +} + +// DNAMEMapping holds a DNAME record's owner and target for redirect synthesis. +type DNAMEMapping struct { + Owner string // e.g., "example.com." + Target string // e.g., "example.net." } func DecodeResponse(msg *dns.Msg) *DecodedResponse { @@ -62,6 +76,7 @@ func DecodeResponse(msg *dns.Msg) *DecodedResponse { Authority: msg.Ns, Additional: msg.Extra, CNAMEChain: extractCNAMEChain(msg), + DNAMEMappings: extractDNAMEMappings(msg), } d.Classification = classify(msg) @@ -75,6 +90,10 @@ func classify(msg *dns.Msg) ResponseClassification { return ResponseNXDOMAIN case dns.RcodeServerFailure: return ResponseSERVFAIL + case dns.RcodeRefused: + return ResponseREFUSED + case dns.RcodeNotImplemented: + return ResponseNOTIMPL case dns.RcodeSuccess: return classifySuccess(msg) default: @@ -125,6 +144,33 @@ func extractCNAMEChain(msg *dns.Msg) []string { return chain } +func extractDNAMEMappings(msg *dns.Msg) []DNAMEMapping { + var mappings []DNAMEMapping + for _, rr := range msg.Answer { + if dname, ok := rr.(*dns.DNAME); ok { + mappings = append(mappings, DNAMEMapping{ + Owner: dns.Fqdn(dname.Hdr.Name), + Target: dns.Fqdn(dname.Target), + }) + } + } + return mappings +} + +// SynthesizeCNAMEFromDNAME computes the CNAME target for queryName given a DNAME mapping. +// Returns empty string if queryName is not a strict subdomain of dnameOwner. +func SynthesizeCNAMEFromDNAME(queryName, dnameOwner, dnameTarget string) string { + q := strings.ToLower(dns.Fqdn(queryName)) + owner := strings.ToLower(dns.Fqdn(dnameOwner)) + target := strings.ToLower(dns.Fqdn(dnameTarget)) + + if !dns.IsSubDomain(owner, q) || q == owner { + return "" + } + prefix := strings.TrimSuffix(q, owner) + return prefix + target +} + func IsTruncated(msg *dns.Msg) bool { return msg != nil && msg.Truncated } diff --git a/internal/dns/query.go b/internal/dns/query.go index 0fb0edc..61b1e82 100644 --- a/internal/dns/query.go +++ b/internal/dns/query.go @@ -93,7 +93,7 @@ func QueryWithExchange(ctx context.Context, server net.IP, name string, qtype ui select { case <-ctx.Done(): return nil, fmt.Errorf("query retries cancelled: %w", ctx.Err()) - case <-time.After(100 * time.Millisecond): + case <-time.After(backoffDelay(attempt)): } } @@ -156,7 +156,7 @@ func IterativeQueryWithExchange(ctx context.Context, server net.IP, name string, select { case <-ctx.Done(): return nil, fmt.Errorf("query retries cancelled: %w", ctx.Err()) - case <-time.After(100 * time.Millisecond): + case <-time.After(backoffDelay(attempt)): } } @@ -195,6 +195,20 @@ func IterativeQueryWithExchange(ctx context.Context, server net.IP, name string, return nil, fmt.Errorf("iterative query %s %s failed after %d retries: %w", name, QNameType(qtype), cfg.Retries, lastErr) } +// backoffDelay computes the wait duration before the given retry attempt (1-indexed). +// Delays: attempt=1 → 100ms, attempt=2 → 200ms, attempt=3 → 400ms, capped at 2s. +func backoffDelay(attempt int) time.Duration { + if attempt <= 0 { + return 0 + } + delay := time.Duration(uint(1)< maxDelay { + return maxDelay + } + return delay +} + func buildQuery(name string, qtype uint16, udpSize int) *dns.Msg { m := new(dns.Msg) m.SetQuestion(dns.Fqdn(name), qtype) diff --git a/internal/dns/robustness_test.go b/internal/dns/robustness_test.go new file mode 100644 index 0000000..c37fc70 --- /dev/null +++ b/internal/dns/robustness_test.go @@ -0,0 +1,139 @@ +package dns + +import ( + "testing" + "time" + + "github.com/miekg/dns" +) + +func TestDecodeResponseREFUSED(t *testing.T) { + msg := newTestMsg(dns.RcodeRefused) + d := DecodeResponse(msg) + if d.Classification != ResponseREFUSED { + t.Errorf("classification = %v, want ResponseREFUSED", d.Classification) + } + if d.RcodeName != "REFUSED" { + t.Errorf("RcodeName = %q, want REFUSED", d.RcodeName) + } +} + +func TestDecodeResponseNOTIMPL(t *testing.T) { + msg := newTestMsg(dns.RcodeNotImplemented) + d := DecodeResponse(msg) + if d.Classification != ResponseNOTIMPL { + t.Errorf("classification = %v, want ResponseNOTIMPL", d.Classification) + } + if d.RcodeName != "NOTIMP" { + t.Errorf("RcodeName = %q, want NOTIMP", d.RcodeName) + } +} + +func TestDecodeResponseREFUSEDString(t *testing.T) { + if got := ResponseREFUSED.String(); got != "refused" { + t.Errorf("ResponseREFUSED.String() = %q, want \"refused\"", got) + } + if got := ResponseNOTIMPL.String(); got != "notimp" { + t.Errorf("ResponseNOTIMPL.String() = %q, want \"notimp\"", got) + } +} + +func TestExtractDNAMEMappings(t *testing.T) { + t.Run("no DNAME", func(t *testing.T) { + msg := new(dns.Msg) + msg.Answer = append(msg.Answer, &dns.A{ + Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA}, + A: MustParseIP("1.2.3.4"), + }) + d := DecodeResponse(msg) + if len(d.DNAMEMappings) != 0 { + t.Errorf("expected 0 DNAME mappings, got %d", len(d.DNAMEMappings)) + } + }) + + t.Run("DNAME in answer", func(t *testing.T) { + msg := new(dns.Msg) + msg.SetReply(new(dns.Msg)) + msg.Answer = append(msg.Answer, &dns.DNAME{ + Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeDNAME, Class: dns.ClassINET, Ttl: 300}, + Target: "example.net.", + }) + d := DecodeResponse(msg) + if len(d.DNAMEMappings) != 1 { + t.Fatalf("expected 1 DNAME mapping, got %d", len(d.DNAMEMappings)) + } + if d.DNAMEMappings[0].Owner != "example.com." { + t.Errorf("Owner = %q, want %q", d.DNAMEMappings[0].Owner, "example.com.") + } + if d.DNAMEMappings[0].Target != "example.net." { + t.Errorf("Target = %q, want %q", d.DNAMEMappings[0].Target, "example.net.") + } + }) +} + +func TestSynthesizeCNAMEFromDNAME(t *testing.T) { + tests := []struct { + queryName string + dnameOwner string + dnameTarget string + want string + }{ + { + queryName: "foo.example.com.", + dnameOwner: "example.com.", + dnameTarget: "example.net.", + want: "foo.example.net.", + }, + { + queryName: "bar.foo.example.com.", + dnameOwner: "example.com.", + dnameTarget: "example.net.", + want: "bar.foo.example.net.", + }, + { + // Owner itself is not redirected + queryName: "example.com.", + dnameOwner: "example.com.", + dnameTarget: "example.net.", + want: "", + }, + { + // Not a subdomain + queryName: "other.com.", + dnameOwner: "example.com.", + dnameTarget: "example.net.", + want: "", + }, + } + + for _, tt := range tests { + got := SynthesizeCNAMEFromDNAME(tt.queryName, tt.dnameOwner, tt.dnameTarget) + if got != tt.want { + t.Errorf("SynthesizeCNAMEFromDNAME(%q, %q, %q) = %q, want %q", + tt.queryName, tt.dnameOwner, tt.dnameTarget, got, tt.want) + } + } +} + +func TestBackoffDelay(t *testing.T) { + tests := []struct { + attempt int + want time.Duration + }{ + {0, 0}, + {1, 100 * time.Millisecond}, + {2, 200 * time.Millisecond}, + {3, 400 * time.Millisecond}, + {4, 800 * time.Millisecond}, + {5, 1600 * time.Millisecond}, + {6, 2000 * time.Millisecond}, // capped at 2s + {10, 2000 * time.Millisecond}, // still capped + } + + for _, tt := range tests { + got := backoffDelay(tt.attempt) + if got != tt.want { + t.Errorf("backoffDelay(%d) = %v, want %v", tt.attempt, got, tt.want) + } + } +} diff --git a/internal/output/stats.go b/internal/output/stats.go index a7df3d2..a6b0456 100644 --- a/internal/output/stats.go +++ b/internal/output/stats.go @@ -138,6 +138,12 @@ func summaryTypeLabel(respType string) string { return "name does not exist" case "servfail": return "resulted in SERVFAIL" + case "refused": + return "query refused by server" + case "notimp": + return "query type not implemented by server" + case "cname_loop": + return "resulted in a CNAME loop" case "error": return "resulted in an error" case "referral": diff --git a/internal/output/text.go b/internal/output/text.go index a5f6fae..dd43c0c 100644 --- a/internal/output/text.go +++ b/internal/output/text.go @@ -210,8 +210,22 @@ func (f *textFormatter) formatResultLine(result traverse.TraversalResult) string return f.colorize(fmt.Sprintf("%s name does not exist", prob), colorYellow) case traverse.RespSERVFAIL: return f.colorize(fmt.Sprintf("%s resulted in SERVFAIL", prob), colorRed) + case traverse.RespREFUSED: + return f.colorize(fmt.Sprintf("%s query refused by server", prob), colorRed) + case traverse.RespNOTIMPL: + return f.colorize(fmt.Sprintf("%s query type not implemented by server", prob), colorRed) + case traverse.RespCNAMELoop: + msg := "CNAME loop detected" + if result.Response.ErrorMessage != "" { + msg = result.Response.ErrorMessage + } + return f.colorize(fmt.Sprintf("%s %s", prob, msg), colorRed) case traverse.RespError: - return f.colorize(fmt.Sprintf("%s resulted in an error", prob), colorRed) + msg := "resulted in an error" + if result.Response.ErrorMessage != "" { + msg = result.Response.ErrorMessage + } + return f.colorize(fmt.Sprintf("%s %s", prob, msg), colorRed) default: return fmt.Sprintf("%s %s", prob, result.Response.Type) } diff --git a/internal/traverse/referral.go b/internal/traverse/referral.go index 3bdfe8a..ca78481 100644 --- a/internal/traverse/referral.go +++ b/internal/traverse/referral.go @@ -8,6 +8,7 @@ import ( "github.com/hits/ExploreDNS/internal/dns" miekgdns "github.com/miekg/dns" + "golang.org/x/net/idna" ) type ResolutionState int @@ -46,12 +47,35 @@ type Referral struct { Prob float64 } +// idnaLookup is the IDN lookup profile used to convert internationalised domain +// names (unicode labels) to their ACE/punycode equivalents before querying. +var idnaLookup = idna.New( + idna.MapForLookup(), + idna.BidiRule(), + idna.StrictDomainName(false), +) + +// toASCII converts a domain name that may contain unicode labels to its +// punycode (ACE) representation. Pure-ASCII names are returned unchanged. +// On conversion errors the original name is returned so the caller can still +// attempt a query (the server will reject it if truly invalid). +func toASCII(name string) string { + if name == "" || name == "." { + return name + } + ascii, err := idnaLookup.ToASCII(name) + if err != nil { + return name + } + return ascii +} + func NewReferral(name string, qtype uint16, bailiwick string, depth int, prob float64, parent *Referral) *Referral { return &Referral{ - Name: miekgdns.Fqdn(strings.ToLower(name)), + Name: miekgdns.Fqdn(strings.ToLower(toASCII(name))), Qtype: qtype, Qclass: miekgdns.ClassINET, - Bailiwick: miekgdns.Fqdn(strings.ToLower(bailiwick)), + Bailiwick: miekgdns.Fqdn(strings.ToLower(toASCII(bailiwick))), Depth: depth, Prob: prob, Parent: parent, @@ -236,3 +260,17 @@ func getVisitedNames(visited map[string]bool) []string { } return names } + +// IsNameInChain reports whether name appears anywhere in this referral's ancestor +// chain, including this referral itself. Used for CNAME loop detection. +func (r *Referral) IsNameInChain(name string) bool { + n := miekgdns.Fqdn(strings.ToLower(name)) + curr := r + for curr != nil { + if curr.Name == n { + return true + } + curr = curr.Parent + } + return false +} diff --git a/internal/traverse/response.go b/internal/traverse/response.go index 0cbc6d6..81620f4 100644 --- a/internal/traverse/response.go +++ b/internal/traverse/response.go @@ -16,6 +16,9 @@ const ( RespNODATA RespNXDOMAIN RespSERVFAIL + RespREFUSED + RespNOTIMPL + RespCNAMELoop RespError ) @@ -33,6 +36,12 @@ func (rt ResponseType) String() string { return "nxdomain" case RespSERVFAIL: return "servfail" + case RespREFUSED: + return "refused" + case RespNOTIMPL: + return "notimp" + case RespCNAMELoop: + return "cname_loop" case RespError: return "error" default: @@ -41,11 +50,12 @@ func (rt ResponseType) String() string { } type Response struct { - Referral *Referral - Server net.IP - Cache *InfoCache - Decoded *dns.DecodedResponse - Type ResponseType + Referral *Referral + Server net.IP + Cache *InfoCache + Decoded *dns.DecodedResponse + Type ResponseType + ErrorMessage string } func NewResponse(ref *Referral, server net.IP, cache *InfoCache) *Response { @@ -59,15 +69,28 @@ func NewResponse(ref *Referral, server net.IP, cache *InfoCache) *Response { func (r *Response) Process(msg *miekgdns.Msg) *Response { if msg == nil { r.Type = RespError + r.ErrorMessage = "nil DNS response" return r } r.Decoded = dns.DecodeResponse(msg) if r.Decoded == nil { r.Type = RespError + r.ErrorMessage = "failed to decode DNS response" return r } + // Synthesize CNAME from DNAME when the server didn't include a synthesized CNAME record. + if len(r.Decoded.CNAMEChain) == 0 && r.Referral != nil && len(r.Decoded.DNAMEMappings) > 0 { + for _, dm := range r.Decoded.DNAMEMappings { + synthesized := dns.SynthesizeCNAMEFromDNAME(r.Referral.Name, dm.Owner, dm.Target) + if synthesized != "" { + r.Decoded.CNAMEChain = append(r.Decoded.CNAMEChain, synthesized) + break + } + } + } + r.Type = r.classify() return r } @@ -78,6 +101,10 @@ func (r *Response) classify() ResponseType { return RespNXDOMAIN case dns.ResponseSERVFAIL: return RespSERVFAIL + case dns.ResponseREFUSED: + return RespREFUSED + case dns.ResponseNOTIMPL: + return RespNOTIMPL case dns.ResponseAnswer: if len(r.Decoded.CNAMEChain) > 0 && !r.hasFinalAnswer() { return RespCNAMEFollow @@ -94,7 +121,11 @@ func (r *Response) classify() ResponseType { func (r *Response) hasFinalAnswer() bool { for _, rr := range r.Decoded.Answers { - if _, ok := rr.(*miekgdns.CNAME); ok { + switch rr.(type) { + case *miekgdns.CNAME, *miekgdns.DNAME, *miekgdns.RRSIG: + // CNAME and DNAME are redirect records, not final answers. + // RRSIG is a DNSSEC signature record — it covers the CNAME/DNAME + // but is not itself the answer to the original question type. continue } return true @@ -200,7 +231,7 @@ func (r *Response) resolveGlue(child *Referral) { func (r *Response) IsTerminal() bool { switch r.Type { - case RespAnswer, RespNODATA, RespNXDOMAIN, RespSERVFAIL, RespError: + case RespAnswer, RespNODATA, RespNXDOMAIN, RespSERVFAIL, RespREFUSED, RespNOTIMPL, RespCNAMELoop, RespError: return true default: return false diff --git a/internal/traverse/robustness_test.go b/internal/traverse/robustness_test.go new file mode 100644 index 0000000..4e1df9a --- /dev/null +++ b/internal/traverse/robustness_test.go @@ -0,0 +1,768 @@ +package traverse + +import ( + "context" + "errors" + "net" + "sync/atomic" + "testing" + + "github.com/miekg/dns" +) + +// TestCNAMELoopDetected verifies that a two-step CNAME loop (A → B → A) is +// detected without infinite recursion and produces a RespCNAMELoop result. +func TestCNAMELoopDetected(t *testing.T) { + // www.example.com → CNAME → alias.example.com → CNAME → www.example.com (loop) + cnameToAlias := new(dns.Msg) + cnameToAlias.SetReply(new(dns.Msg)) + cnameToAlias.Answer = append(cnameToAlias.Answer, &dns.CNAME{ + Hdr: dns.RR_Header{Name: "www.example.com.", Rrtype: dns.TypeCNAME, Class: dns.ClassINET}, + Target: "alias.example.com.", + }) + + cnameBack := new(dns.Msg) + cnameBack.SetReply(new(dns.Msg)) + cnameBack.Answer = append(cnameBack.Answer, &dns.CNAME{ + Hdr: dns.RR_Header{Name: "alias.example.com.", Rrtype: dns.TypeCNAME, Class: dns.ClassINET}, + Target: "www.example.com.", + }) + + tr := NewTraverser(&TraverserConfig{ + MaxDepth: 10, + QueryType: dnsTypeA, + RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, + }) + tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + q := msg.Question[0] + switch q.Name { + case "www.example.com.": + return cnameToAlias.Copy(), nil + case "alias.example.com.": + return cnameBack.Copy(), nil + } + return nil, nil + }) + + ctx := context.Background() + results, err := tr.Traverse(ctx, "www.example.com") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + foundLoop := false + for _, r := range results { + if r.Response != nil && r.Response.Type == RespCNAMELoop { + foundLoop = true + if r.Response.ErrorMessage == "" { + t.Error("expected non-empty ErrorMessage on CNAME loop result") + } + } + } + if !foundLoop { + t.Error("expected RespCNAMELoop result for CNAME loop A → B → A") + } +} + +// TestCNAMEDirectLoop verifies that a direct self-loop (A → A) is handled. +func TestCNAMEDirectLoop(t *testing.T) { + selfLoop := new(dns.Msg) + selfLoop.SetReply(new(dns.Msg)) + selfLoop.Answer = append(selfLoop.Answer, &dns.CNAME{ + Hdr: dns.RR_Header{Name: "www.example.com.", Rrtype: dns.TypeCNAME, Class: dns.ClassINET}, + Target: "www.example.com.", + }) + + tr := NewTraverser(&TraverserConfig{ + MaxDepth: 10, + QueryType: dnsTypeA, + RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, + }) + tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + return selfLoop.Copy(), nil + }) + + ctx := context.Background() + results, err := tr.Traverse(ctx, "www.example.com") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + foundLoop := false + for _, r := range results { + if r.Response != nil && r.Response.Type == RespCNAMELoop { + foundLoop = true + } + } + if !foundLoop { + t.Error("expected RespCNAMELoop for direct self-referencing CNAME") + } +} + +// TestREFUSEDResponse verifies that a REFUSED rcode is classified as RespREFUSED. +func TestREFUSEDResponse(t *testing.T) { + refusedResp := new(dns.Msg) + refusedResp.Rcode = dns.RcodeRefused + + tr := NewTraverser(&TraverserConfig{ + MaxDepth: 5, + QueryType: dnsTypeA, + RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, + }) + tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + return refusedResp.Copy(), nil + }) + + ctx := context.Background() + results, err := tr.Traverse(ctx, "example.com") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if len(results) == 0 { + t.Fatal("expected at least 1 result") + } + if results[0].Response.Type != RespREFUSED { + t.Errorf("Type = %s, want refused", results[0].Response.Type) + } + if !results[0].Response.IsTerminal() { + t.Error("REFUSED should be a terminal response") + } +} + +// TestNOTIMPLResponse verifies that a NOTIMP rcode is classified as RespNOTIMPL. +func TestNOTIMPLResponse(t *testing.T) { + notImplResp := new(dns.Msg) + notImplResp.Rcode = dns.RcodeNotImplemented + + tr := NewTraverser(&TraverserConfig{ + MaxDepth: 5, + QueryType: dnsTypeA, + RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, + }) + tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + return notImplResp.Copy(), nil + }) + + ctx := context.Background() + results, err := tr.Traverse(ctx, "example.com") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if len(results) == 0 { + t.Fatal("expected at least 1 result") + } + if results[0].Response.Type != RespNOTIMPL { + t.Errorf("Type = %s, want notimp", results[0].Response.Type) + } + if !results[0].Response.IsTerminal() { + t.Error("NOTIMP should be a terminal response") + } +} + +// TestGracefulDegradationUnreachableServer verifies that when some servers are +// unreachable, traversal continues with the remaining servers and does not panic. +func TestGracefulDegradationUnreachableServer(t *testing.T) { + answerResp := new(dns.Msg) + answerResp.SetReply(new(dns.Msg)) + answerResp.Answer = append(answerResp.Answer, &dns.A{ + Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, + A: net.ParseIP("93.184.216.34"), + }) + + // Referral with two nameservers; first always fails, second provides the answer. + referralMsg := new(dns.Msg) + referralMsg.Rcode = dns.RcodeSuccess + referralMsg.Authoritative = false + referralMsg.Ns = append(referralMsg.Ns, + &dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns1.example.com."}, + &dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns2.example.com."}, + ) + referralMsg.Extra = append(referralMsg.Extra, + &dns.A{Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.1")}, + &dns.A{Hdr: dns.RR_Header{Name: "ns2.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.2")}, + ) + + tr := NewTraverser(&TraverserConfig{ + MaxDepth: 5, + QueryType: dnsTypeA, + RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, + }) + tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + switch server { + case "198.41.0.4": + return referralMsg.Copy(), nil + case "10.0.0.1": + return nil, errors.New("connection refused") + case "10.0.0.2": + return answerResp.Copy(), nil + } + return nil, errors.New("unexpected server") + }) + + ctx := context.Background() + results, err := tr.Traverse(ctx, "example.com") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + foundAnswer := false + for _, r := range results { + if r.Response != nil && r.Response.Type == RespAnswer { + foundAnswer = true + } + } + if !foundAnswer { + t.Error("expected an answer result from the reachable server") + } +} + +// TestGracefulDegradationAllUnreachable verifies that when ALL servers fail, +// the traversal returns a SERVFAIL result without panicking. +func TestGracefulDegradationAllUnreachable(t *testing.T) { + tr := NewTraverser(&TraverserConfig{ + MaxDepth: 5, + QueryType: dnsTypeA, + RootAddrs: []net.IP{net.ParseIP("198.41.0.4"), net.ParseIP("199.9.14.201")}, + }) + tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + return nil, errors.New("network unreachable") + }) + + ctx := context.Background() + results, err := tr.Traverse(ctx, "example.com") + if err != nil { + t.Fatalf("traversal must not return a top-level error: %v", err) + } + if len(results) == 0 { + t.Fatal("expected at least one result even on total failure") + } + last := results[len(results)-1] + if last.Response == nil { + t.Fatal("last result must have a response") + } + if last.Response.Type != RespSERVFAIL && last.Response.Type != RespError { + t.Errorf("expected SERVFAIL or error when all servers unreachable, got %s", last.Response.Type) + } +} + +// TestDNAMEFollowNoSynthesizedCNAME verifies that a DNAME record in the answer +// section synthesizes a CNAME follow when the server doesn't include one. +func TestDNAMEFollowNoSynthesizedCNAME(t *testing.T) { + // Server returns DNAME only (no synthesized CNAME). + dnameResp := new(dns.Msg) + dnameResp.SetReply(new(dns.Msg)) + dnameResp.Answer = append(dnameResp.Answer, &dns.DNAME{ + Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeDNAME, Class: dns.ClassINET, Ttl: 300}, + Target: "example.net.", + }) + + answerResp := new(dns.Msg) + answerResp.SetReply(new(dns.Msg)) + answerResp.Answer = append(answerResp.Answer, &dns.A{ + Hdr: dns.RR_Header{Name: "www.example.net.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, + A: net.ParseIP("203.0.113.1"), + }) + + tr := NewTraverser(&TraverserConfig{ + MaxDepth: 5, + QueryType: dnsTypeA, + RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, + }) + tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + q := msg.Question[0] + if q.Name == "www.example.com." { + return dnameResp.Copy(), nil + } + if q.Name == "www.example.net." { + return answerResp.Copy(), nil + } + return nil, nil + }) + + ctx := context.Background() + results, err := tr.Traverse(ctx, "www.example.com") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + foundCNAMEFollow := false + for _, r := range results { + if r.Response != nil && r.Response.Type == RespCNAMEFollow { + foundCNAMEFollow = true + } + } + if !foundCNAMEFollow { + t.Error("expected RespCNAMEFollow synthesized from DNAME record") + } +} + +// TestIsNameInChain verifies the ancestor chain lookup. +func TestIsNameInChain(t *testing.T) { + root := NewReferral("example.com", dnsTypeA, ".", 0, 1.0, nil) + child := NewReferral("www.example.com", dnsTypeA, "example.com.", 1, 1.0, root) + grandchild := NewReferral("sub.www.example.com", dnsTypeA, "www.example.com.", 2, 1.0, child) + + tests := []struct { + ref *Referral + name string + want bool + }{ + {grandchild, "sub.www.example.com", true}, // self + {grandchild, "www.example.com", true}, // parent + {grandchild, "example.com", true}, // grandparent + {grandchild, "other.example.com", false}, // not in chain + {root, "example.com", true}, // root matches itself + {root, "www.example.com", false}, // child not in chain from root + } + + for _, tt := range tests { + got := tt.ref.IsNameInChain(tt.name) + if got != tt.want { + t.Errorf("IsNameInChain(%q) from %q = %v, want %v", tt.name, tt.ref.Name, got, tt.want) + } + } +} + +// TestResponseTypeStrings verifies String() for new response types. +func TestResponseTypeStrings(t *testing.T) { + tests := []struct { + rt ResponseType + want string + }{ + {RespReferral, "referral"}, + {RespAnswer, "answer"}, + {RespCNAMEFollow, "cname_follow"}, + {RespNODATA, "nodata"}, + {RespNXDOMAIN, "nxdomain"}, + {RespSERVFAIL, "servfail"}, + {RespREFUSED, "refused"}, + {RespNOTIMPL, "notimp"}, + {RespCNAMELoop, "cname_loop"}, + {RespError, "error"}, + } + + for _, tt := range tests { + if got := tt.rt.String(); got != tt.want { + t.Errorf("ResponseType(%d).String() = %q, want %q", tt.rt, got, tt.want) + } + } +} + +// TestMalformedResponseNoPanic verifies that a nil response from the exchange +// function does not cause a panic, and produces an error result. +func TestMalformedResponseNoPanic(t *testing.T) { + tr := NewTraverser(&TraverserConfig{ + MaxDepth: 5, + QueryType: dnsTypeA, + RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, + }) + tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + return nil, nil // nil response, no error + }) + + ctx := context.Background() + results, err := tr.Traverse(ctx, "example.com") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if len(results) == 0 { + t.Fatal("expected at least one result") + } + // Should produce error/servfail, not panic + for _, r := range results { + if r.Response == nil { + t.Error("result has nil response") + } + } +} + +// TestDNSSECRRSIGDoesNotBlockCNAMEFollow verifies that a DNSSEC RRSIG record +// accompanying a CNAME in the answer section is treated as metadata and does +// NOT prevent the traversal from following the CNAME. +func TestDNSSECRRSIGDoesNotBlockCNAMEFollow(t *testing.T) { + // Server returns CNAME + RRSIG (DNSSEC-signed zone response). + cnameWithRRSIG := new(dns.Msg) + cnameWithRRSIG.SetReply(new(dns.Msg)) + cnameWithRRSIG.Answer = append(cnameWithRRSIG.Answer, + &dns.CNAME{ + Hdr: dns.RR_Header{Name: "www.example.com.", Rrtype: dns.TypeCNAME, Class: dns.ClassINET, Ttl: 300}, + Target: "example.com.", + }, + &dns.RRSIG{ + Hdr: dns.RR_Header{Name: "www.example.com.", Rrtype: dns.TypeRRSIG, Class: dns.ClassINET, Ttl: 300}, + TypeCovered: dns.TypeCNAME, + }, + ) + + finalAnswer := new(dns.Msg) + finalAnswer.SetReply(new(dns.Msg)) + finalAnswer.Answer = append(finalAnswer.Answer, &dns.A{ + Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, + A: net.ParseIP("93.184.216.34"), + }) + + tr := NewTraverser(&TraverserConfig{ + MaxDepth: 10, + QueryType: dnsTypeA, + RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, + Fast: true, + }) + tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + q := msg.Question[0] + if q.Name == "www.example.com." { + return cnameWithRRSIG.Copy(), nil + } + return finalAnswer.Copy(), nil + }) + + ctx := context.Background() + results, err := tr.Traverse(ctx, "www.example.com") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + foundCNAMEFollow := false + foundAnswer := false + for _, r := range results { + if r.Response != nil { + switch r.Response.Type { + case RespCNAMEFollow: + foundCNAMEFollow = true + case RespAnswer: + foundAnswer = true + } + } + } + if !foundCNAMEFollow { + t.Error("expected RespCNAMEFollow: RRSIG should not block CNAME following") + } + if !foundAnswer { + t.Error("expected final RespAnswer after CNAME follow") + } +} + +// TestFastModeOn verifies that Fast=true uses the shared root cache (default +// behaviour): a child branch can see glue stored by the root referral. +func TestFastModeOn(t *testing.T) { + // Root referral returns two nameservers with glue. Each NS branch returns + // an answer. We verify both branches are queried. + referralMsg := new(dns.Msg) + referralMsg.Rcode = dns.RcodeSuccess + referralMsg.Authoritative = false + referralMsg.Ns = append(referralMsg.Ns, + &dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns1.example.com."}, + &dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns2.example.com."}, + ) + referralMsg.Extra = append(referralMsg.Extra, + &dns.A{Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.1")}, + &dns.A{Hdr: dns.RR_Header{Name: "ns2.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.2")}, + ) + + answerMsg := new(dns.Msg) + answerMsg.SetReply(new(dns.Msg)) + answerMsg.Answer = append(answerMsg.Answer, &dns.A{ + Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, + A: net.ParseIP("93.184.216.34"), + }) + + tr := NewTraverser(&TraverserConfig{ + MaxDepth: 5, + QueryType: dnsTypeA, + RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, + Fast: true, + }) + tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + if server == "198.41.0.4" { + return referralMsg.Copy(), nil + } + return answerMsg.Copy(), nil + }) + + ctx := context.Background() + results, err := tr.Traverse(ctx, "example.com") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + answers := 0 + for _, r := range results { + if r.Response != nil && r.Response.Type == RespAnswer { + answers++ + } + } + if answers == 0 { + t.Error("expected at least one answer with Fast=true") + } +} + +// TestFastModeOff verifies that Fast=false gives each referral its own +// independent cache — no cross-branch glue contamination. +func TestFastModeOff(t *testing.T) { + referralMsg := new(dns.Msg) + referralMsg.Rcode = dns.RcodeSuccess + referralMsg.Authoritative = false + referralMsg.Ns = append(referralMsg.Ns, + &dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns1.example.com."}, + &dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns2.example.com."}, + ) + referralMsg.Extra = append(referralMsg.Extra, + &dns.A{Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.1")}, + &dns.A{Hdr: dns.RR_Header{Name: "ns2.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.2")}, + ) + + answerMsg := new(dns.Msg) + answerMsg.SetReply(new(dns.Msg)) + answerMsg.Answer = append(answerMsg.Answer, &dns.A{ + Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, + A: net.ParseIP("93.184.216.34"), + }) + + tr := NewTraverser(&TraverserConfig{ + MaxDepth: 5, + QueryType: dnsTypeA, + RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, + Fast: false, + }) + tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + if server == "198.41.0.4" { + return referralMsg.Copy(), nil + } + return answerMsg.Copy(), nil + }) + + ctx := context.Background() + results, err := tr.Traverse(ctx, "example.com") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + // Traversal must complete without panic and produce results. + if len(results) == 0 { + t.Fatal("expected at least one result with Fast=false") + } +} + +// TestFastModeDefaultIsTrue verifies that DefaultTraverserConfig has Fast=true. +func TestFastModeDefaultIsTrue(t *testing.T) { + cfg := DefaultTraverserConfig() + if !cfg.Fast { + t.Error("DefaultTraverserConfig().Fast should be true") + } +} + +// TestManyNSRecords verifies that a referral with more than 10 nameservers is +// handled gracefully — no panics, results are produced. +func TestManyNSRecords(t *testing.T) { + referralMsg := new(dns.Msg) + referralMsg.Rcode = dns.RcodeSuccess + referralMsg.Authoritative = false + for i := 1; i <= 12; i++ { + ns := &dns.NS{ + Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, + Ns: net.ParseIP(string(rune('a'+i-1))).String() + ".ns.example.com.", + } + // Use a distinct IP for each NS so glue is resolved. + ip := net.IP{10, 0, 0, byte(i)} + referralMsg.Ns = append(referralMsg.Ns, &dns.NS{ + Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, + Ns: ns.Ns, + }) + referralMsg.Extra = append(referralMsg.Extra, &dns.A{ + Hdr: dns.RR_Header{Name: ns.Ns, Rrtype: dnsTypeA}, + A: ip, + }) + } + + answerMsg := new(dns.Msg) + answerMsg.SetReply(new(dns.Msg)) + answerMsg.Answer = append(answerMsg.Answer, &dns.A{ + Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, + A: net.ParseIP("93.184.216.34"), + }) + + var queries int64 + tr := NewTraverser(&TraverserConfig{ + MaxDepth: 5, + QueryType: dnsTypeA, + RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, + Fast: true, + }) + tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + atomic.AddInt64(&queries, 1) + if server == "198.41.0.4" { + return referralMsg.Copy(), nil + } + return answerMsg.Copy(), nil + }) + + ctx := context.Background() + results, err := tr.Traverse(ctx, "example.com") + if err != nil { + t.Fatalf("unexpected error with 12 NS records: %v", err) + } + if len(results) == 0 { + t.Fatal("expected results with many NS records") + } + foundAnswer := false + for _, r := range results { + if r.Response != nil && r.Response.Type == RespAnswer { + foundAnswer = true + } + } + if !foundAnswer { + t.Error("expected at least one answer from the 12-NS referral") + } +} + +// TestIDNPunycodeConversion verifies that a unicode (IDN) domain name is +// converted to its punycode/ACE form before querying. +func TestIDNPunycodeConversion(t *testing.T) { + // "münchen.de" → "xn--mnchen-3ya.de" (after punycode encoding) + ref := NewReferral("münchen.de", dnsTypeA, ".", 0, 1.0, nil) + if ref.Name == "münchen.de." { + t.Errorf("IDN name was not converted to punycode: got %q", ref.Name) + } + // Verify it starts with the expected punycode label. + if ref.Name != "xn--mnchen-3ya.de." { + t.Errorf("unexpected punycode result: got %q, want %q", ref.Name, "xn--mnchen-3ya.de.") + } +} + +// TestASCIIDomainUnchanged verifies that a plain ASCII domain is not mangled +// by the IDN conversion path. +func TestASCIIDomainUnchanged(t *testing.T) { + ref := NewReferral("example.com", dnsTypeA, ".", 0, 1.0, nil) + if ref.Name != "example.com." { + t.Errorf("ASCII domain was mangled: got %q, want %q", ref.Name, "example.com.") + } +} + +// TestWildcardResponse verifies that a wildcard answer (e.g. *.example.com +// returning an A record for sub.example.com) is handled as a regular answer. +func TestWildcardResponse(t *testing.T) { + wildcardAnswer := new(dns.Msg) + wildcardAnswer.SetReply(new(dns.Msg)) + wildcardAnswer.Authoritative = true + wildcardAnswer.Answer = append(wildcardAnswer.Answer, &dns.A{ + Hdr: dns.RR_Header{Name: "sub.example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, + A: net.ParseIP("1.2.3.4"), + }) + + tr := NewTraverser(&TraverserConfig{ + MaxDepth: 5, + QueryType: dnsTypeA, + RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, + }) + tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + return wildcardAnswer.Copy(), nil + }) + + ctx := context.Background() + results, err := tr.Traverse(ctx, "sub.example.com") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if len(results) == 0 { + t.Fatal("expected results for wildcard response") + } + if results[0].Response.Type != RespAnswer { + t.Errorf("Type = %s, want answer", results[0].Response.Type) + } +} + +// TestLongCNAMEChainDepthLimit verifies that a very long CNAME chain is +// terminated by the MaxDepth limit without infinite recursion or a panic. +func TestLongCNAMEChainDepthLimit(t *testing.T) { + // Every query returns a CNAME to the next label. The MaxDepth setting + // must stop the chain. + counter := 0 + tr := NewTraverser(&TraverserConfig{ + MaxDepth: 5, + QueryType: dnsTypeA, + RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, + }) + tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + counter++ + q := msg.Question[0] + resp := new(dns.Msg) + resp.SetReply(msg) + next := "next" + q.Name + resp.Answer = append(resp.Answer, &dns.CNAME{ + Hdr: dns.RR_Header{Name: q.Name, Rrtype: dnsTypeCNAME, Class: dns.ClassINET}, + Target: next, + }) + return resp, nil + }) + + ctx := context.Background() + results, err := tr.Traverse(ctx, "start.example.com") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if len(results) == 0 { + t.Fatal("expected results") + } + // Traversal must have stopped — counter should not be unbounded. + if counter > 50 { + t.Errorf("too many exchange calls (%d): chain depth limit not enforced", counter) + } +} + +// TestPartialBranchFailureReturnsResults verifies the graceful degradation +// requirement: when some NS branches fail completely, the partial results from +// successful branches are still returned. +func TestPartialBranchFailureReturnsResults(t *testing.T) { + // Three nameservers: first two error, third succeeds. + referralMsg := new(dns.Msg) + referralMsg.Rcode = dns.RcodeSuccess + referralMsg.Authoritative = false + referralMsg.Ns = append(referralMsg.Ns, + &dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns1.example.com."}, + &dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns2.example.com."}, + &dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns3.example.com."}, + ) + referralMsg.Extra = append(referralMsg.Extra, + &dns.A{Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.1")}, + &dns.A{Hdr: dns.RR_Header{Name: "ns2.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.2")}, + &dns.A{Hdr: dns.RR_Header{Name: "ns3.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.3")}, + ) + + answerMsg := new(dns.Msg) + answerMsg.SetReply(new(dns.Msg)) + answerMsg.Answer = append(answerMsg.Answer, &dns.A{ + Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, + A: net.ParseIP("93.184.216.34"), + }) + + tr := NewTraverser(&TraverserConfig{ + MaxDepth: 5, + QueryType: dnsTypeA, + RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, + }) + tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + switch server { + case "198.41.0.4": + return referralMsg.Copy(), nil + case "10.0.0.1", "10.0.0.2": + return nil, errors.New("server unreachable") + case "10.0.0.3": + return answerMsg.Copy(), nil + } + return nil, errors.New("unexpected server") + }) + + ctx := context.Background() + results, err := tr.Traverse(ctx, "example.com") + if err != nil { + t.Fatalf("traversal must not return a top-level error: %v", err) + } + + foundAnswer := false + for _, r := range results { + if r.Response != nil && r.Response.Type == RespAnswer { + foundAnswer = true + } + } + if !foundAnswer { + t.Error("expected an answer from the third (reachable) nameserver despite others failing") + } +} diff --git a/internal/traverse/traverser.go b/internal/traverse/traverser.go index 3b2c259..c09f835 100644 --- a/internal/traverse/traverser.go +++ b/internal/traverse/traverser.go @@ -5,6 +5,7 @@ import ( "fmt" "net" "sync" + "time" "github.com/hits/ExploreDNS/internal/dns" miekgdns "github.com/miekg/dns" @@ -17,6 +18,12 @@ type TraverserConfig struct { QueryConfig *dns.QueryConfig RootAddrs []net.IP 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 { @@ -26,6 +33,7 @@ func DefaultTraverserConfig() *TraverserConfig { RootConfig: nil, QueryConfig: nil, RootAddrs: nil, + Fast: true, } } @@ -97,9 +105,19 @@ func (t *Traverser) Traverse(ctx context.Context, name string) ([]TraversalResul break } - cache := rootCache - if ref.Parent != nil { - cache = rootCache.Child() + 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 { @@ -141,7 +159,19 @@ func (t *Traverser) Traverse(ctx context.Context, name string) ([]TraversalResul if resp.Type == RespCNAMEFollow { follow := resp.CNAMEFollowReferral() if follow != nil { - if !stack.Push(follow) { + // 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, @@ -412,12 +442,16 @@ func (t *Traverser) resolveGlueViaSystem(ctx context.Context, name string, cache c := &miekgdns.Client{ Net: "udp", - ReadTimeout: 5, - WriteTimeout: 5, + ReadTimeout: 5 * time.Second, + WriteTimeout: 5 * time.Second, } if deadline, ok := ctx.Deadline(); ok { - c.ReadTimeout = deadline.Sub(deadline) - c.WriteTimeout = deadline.Sub(deadline) + remaining := time.Until(deadline) + if remaining <= 0 { + return nil + } + c.ReadTimeout = remaining + c.WriteTimeout = remaining } fqdn := miekgdns.Fqdn(name)