Files
ExploreDNS/internal/traverse/robustness_test.go
2026-06-07 17:24:00 +00:00

769 lines
24 KiB
Go

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")
}
}