Files
ExploreDNS/internal/traverse/traverser_test.go
T
Garyandmultica-agent 2835ee86cf
CI / test (pull_request) Failing after 2m40s
feat: implement Phase 2.3 DNS traversal engine
Add core traversal engine with iterative resolution from root servers:
- referral.go: Referral struct with name, type, bailiwick, addresses, state
- cache.go: Chained InfoCache with parent inheritance for NS and glue records
- stack.go: LIFO stack with configurable max depth enforcement
- response.go: Response classifier (referral, answer, CNAME follow, NODATA,
  NXDOMAIN, SERVFAIL) with child referral generation and probability splitting
- traverser.go: Traverser orchestrates full traversal from root to leaf,
  following all referral branches, handling CNAME chains, and respecting
  max depth limits

Add IterativeQuery/IterativeQueryWithExchange to dns package (RD=false queries).

57 new tests covering all acceptance criteria: traversal from root to leaf,
comprehensive branch following, max depth enforcement, probability distribution,
chained cache behavior, CNAME following, NXDOMAIN/SERVFAIL handling, context
cancellation, and concurrent access.

Co-authored-by: multica-agent <github@multica.ai>
2026-06-06 13:23:01 +10:00

484 lines
13 KiB
Go

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