Co-authored-by: Hansen IT Solutions <gary@hansenits.com> Co-committed-by: Hansen IT Solutions <gary@hansenits.com>
This commit was merged in pull request #6.
This commit is contained in:
@@ -0,0 +1,483 @@
|
||||
package traverse
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"testing"
|
||||
|
||||
"github.com/miekg/dns"
|
||||
)
|
||||
|
||||
const (
|
||||
dnsTypeA = dns.TypeA
|
||||
dnsTypeNS = dns.TypeNS
|
||||
dnsTypeCNAME = dns.TypeCNAME
|
||||
dnsTypeSOA = dns.TypeSOA
|
||||
)
|
||||
|
||||
func TestDefaultTraverserConfig(t *testing.T) {
|
||||
cfg := DefaultTraverserConfig()
|
||||
if cfg.MaxDepth != DefaultMaxDepth {
|
||||
t.Errorf("MaxDepth = %d, want %d", cfg.MaxDepth, DefaultMaxDepth)
|
||||
}
|
||||
if cfg.QueryType != dnsTypeA {
|
||||
t.Errorf("QueryType = %d, want %d", cfg.QueryType, dnsTypeA)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewTraverserNilConfig(t *testing.T) {
|
||||
tr := NewTraverser(nil)
|
||||
if tr == nil {
|
||||
t.Fatal("NewTraverser(nil) should not return nil")
|
||||
}
|
||||
}
|
||||
|
||||
func TestTraverserSimpleTraversal(t *testing.T) {
|
||||
answerResp := func() *dns.Msg {
|
||||
m := new(dns.Msg)
|
||||
m.SetReply(new(dns.Msg))
|
||||
m.Answer = append(m.Answer, &dns.A{
|
||||
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300},
|
||||
A: net.ParseIP("93.184.216.34"),
|
||||
})
|
||||
return m
|
||||
}()
|
||||
|
||||
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 answerResp.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")
|
||||
}
|
||||
|
||||
found := false
|
||||
for _, r := range results {
|
||||
if r.Response.Type == RespAnswer {
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
t.Error("expected to find an answer response")
|
||||
}
|
||||
}
|
||||
|
||||
func TestTraverserReferralTraversal(t *testing.T) {
|
||||
rootAnswer := new(dns.Msg)
|
||||
rootAnswer.Rcode = dns.RcodeSuccess
|
||||
rootAnswer.Authoritative = false
|
||||
rootAnswer.Ns = append(rootAnswer.Ns, &dns.NS{
|
||||
Hdr: dns.RR_Header{Name: "com.", Rrtype: dnsTypeNS, Class: dns.ClassINET},
|
||||
Ns: "a.gtld-servers.net.",
|
||||
})
|
||||
rootAnswer.Extra = append(rootAnswer.Extra, &dns.A{
|
||||
Hdr: dns.RR_Header{Name: "a.gtld-servers.net.", Rrtype: dnsTypeA},
|
||||
A: net.ParseIP("192.5.6.30"),
|
||||
})
|
||||
|
||||
tldAnswer := new(dns.Msg)
|
||||
tldAnswer.SetReply(new(dns.Msg))
|
||||
tldAnswer.Answer = append(tldAnswer.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) {
|
||||
q := msg.Question[0]
|
||||
key := q.Name + "/" + dns.TypeToString[q.Qtype]
|
||||
if q.Name == "example.com." && server == "198.41.0.4" {
|
||||
return rootAnswer.Copy(), nil
|
||||
}
|
||||
if q.Name == "example.com." {
|
||||
return tldAnswer.Copy(), nil
|
||||
}
|
||||
_ = key
|
||||
return nil, nil
|
||||
})
|
||||
|
||||
ctx := context.Background()
|
||||
results, err := tr.Traverse(ctx, "example.com")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if len(results) < 2 {
|
||||
t.Fatalf("expected at least 2 results (referral + answer), got %d", len(results))
|
||||
}
|
||||
}
|
||||
|
||||
func TestTraverserMaxDepth(t *testing.T) {
|
||||
callCount := 0
|
||||
tr := NewTraverser(&TraverserConfig{
|
||||
MaxDepth: 2,
|
||||
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) {
|
||||
callCount++
|
||||
m := new(dns.Msg)
|
||||
m.Rcode = dns.RcodeSuccess
|
||||
m.Authoritative = false
|
||||
m.Ns = append(m.Ns, &dns.NS{
|
||||
Hdr: dns.RR_Header{Rrtype: dnsTypeNS},
|
||||
Ns: "ns.example.com.",
|
||||
})
|
||||
m.Extra = append(m.Extra, &dns.A{
|
||||
Hdr: dns.RR_Header{Name: "ns.example.com.", Rrtype: dnsTypeA},
|
||||
A: net.ParseIP("1.2.3.4"),
|
||||
})
|
||||
return m, nil
|
||||
})
|
||||
|
||||
ctx := context.Background()
|
||||
results, err := tr.Traverse(ctx, "deep.example.com")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
|
||||
if callCount < 1 {
|
||||
t.Errorf("expected at least 1 call before max depth, got %d", callCount)
|
||||
}
|
||||
|
||||
depthExceeded := false
|
||||
for _, r := range results {
|
||||
if r.Referral != nil && r.Referral.Depth >= 2 {
|
||||
depthExceeded = true
|
||||
}
|
||||
if r.Response.Type == RespError {
|
||||
depthExceeded = true
|
||||
}
|
||||
}
|
||||
if !depthExceeded {
|
||||
t.Error("expected to see depth exceeded results")
|
||||
}
|
||||
}
|
||||
|
||||
func TestTraverserContextCancellation(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) {
|
||||
m := new(dns.Msg)
|
||||
m.Rcode = dns.RcodeSuccess
|
||||
m.Authoritative = false
|
||||
m.Ns = append(m.Ns, &dns.NS{
|
||||
Hdr: dns.RR_Header{Rrtype: dnsTypeNS},
|
||||
Ns: "ns.example.com.",
|
||||
})
|
||||
m.Extra = append(m.Extra, &dns.A{
|
||||
Hdr: dns.RR_Header{Name: "ns.example.com.", Rrtype: dnsTypeA},
|
||||
A: net.ParseIP("1.2.3.4"),
|
||||
})
|
||||
return m, nil
|
||||
})
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
|
||||
_, err := tr.Traverse(ctx, "example.com")
|
||||
if err == nil {
|
||||
t.Fatal("expected error on cancelled context")
|
||||
}
|
||||
}
|
||||
|
||||
func TestTraverserNXDOMAIN(t *testing.T) {
|
||||
nxdResp := new(dns.Msg)
|
||||
nxdResp.Rcode = dns.RcodeNameError
|
||||
|
||||
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 nxdResp.Copy(), nil
|
||||
})
|
||||
|
||||
ctx := context.Background()
|
||||
results, err := tr.Traverse(ctx, "nonexistent.invalid")
|
||||
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 != RespNXDOMAIN {
|
||||
t.Errorf("Type = %d, want %d", results[0].Response.Type, RespNXDOMAIN)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTraverserSERVFAIL(t *testing.T) {
|
||||
sfResp := new(dns.Msg)
|
||||
sfResp.Rcode = dns.RcodeServerFailure
|
||||
|
||||
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 sfResp.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 != RespSERVFAIL {
|
||||
t.Errorf("Type = %d, want %d", results[0].Response.Type, RespSERVFAIL)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTraverserCNAMEFollow(t *testing.T) {
|
||||
cnameResp := new(dns.Msg)
|
||||
cnameResp.SetReply(new(dns.Msg))
|
||||
cnameResp.Answer = append(cnameResp.Answer,
|
||||
&dns.CNAME{
|
||||
Hdr: dns.RR_Header{Name: "www.example.com.", Rrtype: dnsTypeCNAME, Class: dns.ClassINET},
|
||||
Target: "example.com.",
|
||||
},
|
||||
)
|
||||
|
||||
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"),
|
||||
})
|
||||
|
||||
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 cnameResp.Copy(), nil
|
||||
}
|
||||
if q.Name == "example.com." {
|
||||
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)
|
||||
}
|
||||
|
||||
foundCNAME := false
|
||||
foundAnswer := false
|
||||
for _, r := range results {
|
||||
if r.Response.Type == RespCNAMEFollow {
|
||||
foundCNAME = true
|
||||
}
|
||||
if r.Response.Type == RespAnswer {
|
||||
foundAnswer = true
|
||||
}
|
||||
}
|
||||
if !foundCNAME {
|
||||
t.Error("expected CNAME follow response")
|
||||
}
|
||||
if !foundAnswer {
|
||||
t.Error("expected final answer response")
|
||||
}
|
||||
}
|
||||
|
||||
func TestTraverserProbabilityCalculation(t *testing.T) {
|
||||
rootReferral := new(dns.Msg)
|
||||
rootReferral.Rcode = dns.RcodeSuccess
|
||||
rootReferral.Authoritative = false
|
||||
rootReferral.Ns = append(rootReferral.Ns,
|
||||
&dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dnsTypeNS}, Ns: "a.root-servers.net."},
|
||||
&dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dnsTypeNS}, Ns: "b.root-servers.net."},
|
||||
)
|
||||
rootReferral.Extra = append(rootReferral.Extra,
|
||||
&dns.A{Hdr: dns.RR_Header{Name: "a.root-servers.net.", Rrtype: dnsTypeA}, A: net.ParseIP("198.41.0.4")},
|
||||
&dns.A{Hdr: dns.RR_Header{Name: "b.root-servers.net.", Rrtype: dnsTypeA}, A: net.ParseIP("199.9.14.201")},
|
||||
)
|
||||
|
||||
tldAnswer := new(dns.Msg)
|
||||
tldAnswer.SetReply(new(dns.Msg))
|
||||
tldAnswer.Answer = append(tldAnswer.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("1.2.3.4")},
|
||||
})
|
||||
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||
q := msg.Question[0]
|
||||
if q.Name == "example.com." && server == "1.2.3.4" {
|
||||
return rootReferral.Copy(), nil
|
||||
}
|
||||
return tldAnswer.Copy(), nil
|
||||
})
|
||||
|
||||
ctx := context.Background()
|
||||
results, err := tr.Traverse(ctx, "example.com")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
|
||||
for _, r := range results {
|
||||
if r.Referral != nil && r.Referral.Depth == 1 && r.Referral.Parent != nil {
|
||||
if r.Referral.Prob != 0.5 {
|
||||
t.Errorf("child prob = %f, want 0.5", r.Referral.Prob)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestTraverserNODATA(t *testing.T) {
|
||||
nodataResp := new(dns.Msg)
|
||||
nodataResp.Rcode = dns.RcodeSuccess
|
||||
nodataResp.Authoritative = true
|
||||
nodataResp.Ns = append(nodataResp.Ns, &dns.SOA{
|
||||
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeSOA, Class: dns.ClassINET, Ttl: 3600},
|
||||
})
|
||||
|
||||
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 nodataResp.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 != RespNODATA {
|
||||
t.Errorf("Type = %d, want %d", results[0].Response.Type, RespNODATA)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTraverserMultipleRoots(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"),
|
||||
})
|
||||
|
||||
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 answerResp.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 results")
|
||||
}
|
||||
}
|
||||
|
||||
func TestTraverserNilExchangeResponse(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
|
||||
})
|
||||
|
||||
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 even with nil response")
|
||||
}
|
||||
}
|
||||
|
||||
func TestTraverserCacheChaining(t *testing.T) {
|
||||
rootReferral := new(dns.Msg)
|
||||
rootReferral.Rcode = dns.RcodeSuccess
|
||||
rootReferral.Authoritative = false
|
||||
rootReferral.Ns = append(rootReferral.Ns,
|
||||
&dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dnsTypeNS}, Ns: "a.gtld-servers.net."},
|
||||
)
|
||||
rootReferral.Extra = append(rootReferral.Extra,
|
||||
&dns.A{Hdr: dns.RR_Header{Name: "a.gtld-servers.net.", Rrtype: dnsTypeA}, A: net.ParseIP("192.5.6.30")},
|
||||
)
|
||||
|
||||
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"),
|
||||
})
|
||||
|
||||
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 == "example.com." && server == "198.41.0.4" {
|
||||
return rootReferral.Copy(), nil
|
||||
}
|
||||
return answerResp.Copy(), nil
|
||||
})
|
||||
|
||||
ctx := context.Background()
|
||||
results, err := tr.Traverse(ctx, "example.com")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
|
||||
cacheHits := 0
|
||||
for _, r := range results {
|
||||
if r.Response != nil && r.Response.Cache != nil {
|
||||
if r.Response.Cache.NSCount() > 0 {
|
||||
cacheHits++
|
||||
}
|
||||
}
|
||||
}
|
||||
if cacheHits == 0 {
|
||||
t.Error("expected cache to store NS records from referrals")
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user