- Update go.mod module declaration - Update all internal import paths in .go files - Update go install lines in README.md Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Co-authored-by: multica-agent <github@multica.ai>
728 lines
21 KiB
Go
728 lines
21 KiB
Go
package traverse
|
|
|
|
import (
|
|
"context"
|
|
"net"
|
|
"testing"
|
|
"time"
|
|
|
|
dnsinternal "gitea.hansenits.com.au/hits/ExploreDNS/internal/dns"
|
|
"github.com/miekg/dns"
|
|
)
|
|
|
|
func TestSetHooks(t *testing.T) {
|
|
tr := NewTraverser(nil)
|
|
hooks := &TraverserHooks{
|
|
OnEvent: func(event TraversalEvent) {},
|
|
}
|
|
tr.SetHooks(hooks)
|
|
if tr.config.Hooks != hooks {
|
|
t.Error("SetHooks should set config.Hooks")
|
|
}
|
|
|
|
// SetHooks on nil config traverser (initializes config)
|
|
tr2 := &Traverser{}
|
|
tr2.SetHooks(hooks)
|
|
if tr2.config == nil || tr2.config.Hooks != hooks {
|
|
t.Error("SetHooks should initialize config when nil")
|
|
}
|
|
}
|
|
|
|
func TestNewAQuery(t *testing.T) {
|
|
msg := newAQuery("example.com.")
|
|
if msg == nil {
|
|
t.Fatal("newAQuery returned nil")
|
|
}
|
|
if !msg.RecursionDesired {
|
|
t.Error("expected RD=true in newAQuery")
|
|
}
|
|
if len(msg.Question) == 0 {
|
|
t.Fatal("expected question in newAQuery")
|
|
}
|
|
if msg.Question[0].Qtype != dns.TypeA {
|
|
t.Errorf("expected TypeA, got %d", msg.Question[0].Qtype)
|
|
}
|
|
}
|
|
|
|
func TestResolveGlueViaSystemCacheHit(t *testing.T) {
|
|
tr := NewTraverser(&TraverserConfig{
|
|
MaxDepth: 5,
|
|
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
|
})
|
|
cache := NewInfoCache(nil)
|
|
expected := []net.IP{net.ParseIP("1.2.3.4")}
|
|
cache.StoreGlue("ns1.example.com.", expected)
|
|
|
|
ctx := context.Background()
|
|
addrs := tr.resolveGlueViaSystem(ctx, "ns1.example.com.", cache)
|
|
if len(addrs) == 0 {
|
|
t.Error("expected addresses from cache hit")
|
|
}
|
|
}
|
|
|
|
func TestResolveGlueViaSystemExpiredContext(t *testing.T) {
|
|
tr := NewTraverser(&TraverserConfig{
|
|
MaxDepth: 5,
|
|
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
|
})
|
|
|
|
// Expired context → remaining <= 0 → returns nil immediately
|
|
ctx, cancel := context.WithDeadline(context.Background(), time.Now().Add(-time.Second))
|
|
defer cancel()
|
|
|
|
addrs := tr.resolveGlueViaSystem(ctx, "ns1.example.com.", nil)
|
|
if len(addrs) != 0 {
|
|
t.Errorf("expected nil from expired context, got %v", addrs)
|
|
}
|
|
}
|
|
|
|
func TestResolveGlueViaSystemTimeout(t *testing.T) {
|
|
tr := NewTraverser(&TraverserConfig{
|
|
MaxDepth: 5,
|
|
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
|
})
|
|
|
|
// Very short timeout will fail the DNS query to 127.0.0.1:53
|
|
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Millisecond)
|
|
defer cancel()
|
|
time.Sleep(15 * time.Millisecond) // ensure it's expired
|
|
|
|
addrs := tr.resolveGlueViaSystem(ctx, "ns1.example.com.", nil)
|
|
// May return nil (timeout) or addresses (if local resolver responds instantly)
|
|
t.Logf("resolveGlueViaSystem returned %d addresses", len(addrs))
|
|
}
|
|
|
|
func TestEnsureRDFalseWithExchange(t *testing.T) {
|
|
rdFalseMsg := new(dns.Msg)
|
|
rdFalseMsg.SetReply(new(dns.Msg))
|
|
rdFalseMsg.RecursionDesired = false
|
|
|
|
tr := NewTraverser(&TraverserConfig{
|
|
MaxDepth: 5,
|
|
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 rdFalseMsg.Copy(), nil
|
|
})
|
|
|
|
rdTrueMsg := new(dns.Msg)
|
|
rdTrueMsg.SetReply(new(dns.Msg))
|
|
rdTrueMsg.RecursionDesired = true
|
|
|
|
result := tr.ensureRDFalse(rdTrueMsg, net.ParseIP("198.41.0.4"), "example.com.", dnsinternal.TypeA, nil)
|
|
if result == nil {
|
|
t.Fatal("ensureRDFalse with exchange should return non-nil")
|
|
}
|
|
}
|
|
|
|
func TestEnsureRDFalseNilMsg(t *testing.T) {
|
|
tr := NewTraverser(&TraverserConfig{
|
|
MaxDepth: 5,
|
|
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
|
})
|
|
result := tr.ensureRDFalse(nil, net.ParseIP("1.2.3.4"), "example.com.", dnsinternal.TypeA, nil)
|
|
if result != nil {
|
|
t.Error("ensureRDFalse(nil) should return nil")
|
|
}
|
|
}
|
|
|
|
func TestEnsureRDFalseRDAlreadyFalse(t *testing.T) {
|
|
tr := NewTraverser(&TraverserConfig{
|
|
MaxDepth: 5,
|
|
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
|
})
|
|
|
|
msg := new(dns.Msg)
|
|
msg.RecursionDesired = false
|
|
result := tr.ensureRDFalse(msg, net.ParseIP("1.2.3.4"), "example.com.", dnsinternal.TypeA, nil)
|
|
if result != msg {
|
|
t.Error("ensureRDFalse should return same msg when RD=false")
|
|
}
|
|
}
|
|
|
|
func TestEnsureRDFalseNoExchange(t *testing.T) {
|
|
tr := NewTraverser(&TraverserConfig{
|
|
MaxDepth: 5,
|
|
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
|
})
|
|
// No exchange set
|
|
|
|
msg := new(dns.Msg)
|
|
msg.RecursionDesired = true
|
|
result := tr.ensureRDFalse(msg, net.ParseIP("1.2.3.4"), "example.com.", dnsinternal.TypeA, nil)
|
|
if result == nil {
|
|
t.Fatal("ensureRDFalse without exchange should return msg with RD cleared")
|
|
}
|
|
if result.RecursionDesired {
|
|
t.Error("expected RD=false after ensureRDFalse without exchange")
|
|
}
|
|
}
|
|
|
|
func TestResolveNSFromCache(t *testing.T) {
|
|
tr := NewTraverser(&TraverserConfig{
|
|
MaxDepth: 5,
|
|
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
|
|
})
|
|
|
|
cache := NewInfoCache(nil)
|
|
expected := []net.IP{net.ParseIP("1.2.3.4")}
|
|
cache.StoreGlue("ns1.example.com.", expected)
|
|
|
|
ctx := context.Background()
|
|
addrs, err := tr.ResolveNS(ctx, "ns1.example.com.", cache, nil, 0)
|
|
if err != nil {
|
|
t.Fatalf("ResolveNS cache hit: %v", err)
|
|
}
|
|
if len(addrs) == 0 {
|
|
t.Error("expected addresses from cache")
|
|
}
|
|
}
|
|
|
|
func TestResolveNSCircularReferral(t *testing.T) {
|
|
tr := NewTraverser(&TraverserConfig{
|
|
MaxDepth: 5,
|
|
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
|
|
})
|
|
|
|
visited := map[string]bool{"ns1.example.com.": true}
|
|
ctx := context.Background()
|
|
_, err := tr.ResolveNS(ctx, "ns1.example.com.", nil, visited, 0)
|
|
if err == nil {
|
|
t.Fatal("expected circular referral error")
|
|
}
|
|
var circErr *CircularReferralError
|
|
if _, ok := err.(*CircularReferralError); !ok {
|
|
t.Errorf("expected CircularReferralError, got %T: %v", err, err)
|
|
}
|
|
_ = circErr
|
|
}
|
|
|
|
func TestResolveNSMaxDepth(t *testing.T) {
|
|
tr := NewTraverser(&TraverserConfig{
|
|
MaxDepth: 5,
|
|
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()
|
|
_, err := tr.ResolveNS(ctx, "ns1.example.com.", nil, nil, DefaultMaxDepth+1)
|
|
if err == nil {
|
|
t.Fatal("expected max depth error")
|
|
}
|
|
if _, ok := err.(*UnresolvableNameserverError); !ok {
|
|
t.Errorf("expected UnresolvableNameserverError, got %T: %v", err, err)
|
|
}
|
|
}
|
|
|
|
func TestResolveNSWithAnswer(t *testing.T) {
|
|
answerMsg := new(dns.Msg)
|
|
answerMsg.SetReply(new(dns.Msg))
|
|
answerMsg.Answer = append(answerMsg.Answer, &dns.A{
|
|
Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300},
|
|
A: net.ParseIP("1.2.3.4"),
|
|
})
|
|
|
|
tr := NewTraverser(&TraverserConfig{
|
|
MaxDepth: 5,
|
|
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 answerMsg.Copy(), nil
|
|
})
|
|
|
|
ctx := context.Background()
|
|
addrs, err := tr.ResolveNS(ctx, "ns1.example.com.", nil, nil, 0)
|
|
if err != nil {
|
|
t.Fatalf("ResolveNS with answer: %v", err)
|
|
}
|
|
if len(addrs) == 0 {
|
|
t.Fatal("expected addresses from NS resolution")
|
|
}
|
|
}
|
|
|
|
func TestResolveNSNXDOMAIN(t *testing.T) {
|
|
nxMsg := new(dns.Msg)
|
|
nxMsg.Rcode = dns.RcodeNameError
|
|
|
|
tr := NewTraverser(&TraverserConfig{
|
|
MaxDepth: 5,
|
|
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 nxMsg.Copy(), nil
|
|
})
|
|
|
|
ctx := context.Background()
|
|
_, err := tr.ResolveNS(ctx, "nonexistent.invalid.", nil, nil, 0)
|
|
if err == nil {
|
|
t.Fatal("expected error for NXDOMAIN NS resolution")
|
|
}
|
|
}
|
|
|
|
func TestResolveNSContextCancellation(t *testing.T) {
|
|
tr := NewTraverser(&TraverserConfig{
|
|
MaxDepth: 5,
|
|
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) {
|
|
// Keep returning referrals to keep the loop going
|
|
refMsg := new(dns.Msg)
|
|
refMsg.Rcode = dns.RcodeSuccess
|
|
refMsg.Ns = append(refMsg.Ns, &dns.NS{
|
|
Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS, Class: dns.ClassINET},
|
|
Ns: "ns.example.com.",
|
|
})
|
|
refMsg.Extra = append(refMsg.Extra, &dns.A{
|
|
Hdr: dns.RR_Header{Name: "ns.example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET},
|
|
A: net.ParseIP("1.2.3.4"),
|
|
})
|
|
return refMsg, nil
|
|
})
|
|
|
|
ctx, cancel := context.WithCancel(context.Background())
|
|
cancel() // Cancel immediately
|
|
|
|
_, err := tr.ResolveNS(ctx, "ns1.example.com.", nil, nil, 0)
|
|
if err == nil {
|
|
t.Fatal("expected error on cancelled context")
|
|
}
|
|
}
|
|
|
|
func TestDiscoverRootsWithRootAddrs(t *testing.T) {
|
|
expected := []net.IP{net.ParseIP("198.41.0.4"), net.ParseIP("199.9.14.201")}
|
|
tr := NewTraverser(&TraverserConfig{
|
|
MaxDepth: 5,
|
|
RootAddrs: expected,
|
|
})
|
|
|
|
ctx := context.Background()
|
|
addrs, err := tr.discoverRoots(ctx)
|
|
if err != nil {
|
|
t.Fatalf("discoverRoots with RootAddrs: %v", err)
|
|
}
|
|
if len(addrs) != len(expected) {
|
|
t.Errorf("expected %d addresses, got %d", len(expected), len(addrs))
|
|
}
|
|
}
|
|
|
|
func TestDiscoverRootsFromSystem(t *testing.T) {
|
|
tr := NewTraverser(&TraverserConfig{
|
|
MaxDepth: 5,
|
|
// No RootAddrs - will call dns.DiscoverRoots
|
|
})
|
|
|
|
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
|
defer cancel()
|
|
|
|
addrs, err := tr.discoverRoots(ctx)
|
|
if err != nil {
|
|
t.Logf("discoverRoots without RootAddrs error (may skip): %v", err)
|
|
t.Skip()
|
|
}
|
|
if len(addrs) == 0 {
|
|
t.Error("expected at least one root address")
|
|
}
|
|
}
|
|
|
|
func TestTraverserSetHooksAndTraverse(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: dns.TypeA, Class: dns.ClassINET, Ttl: 300},
|
|
A: net.ParseIP("1.2.3.4"),
|
|
})
|
|
|
|
tr := NewTraverser(&TraverserConfig{
|
|
MaxDepth: 5,
|
|
QueryType: dnsinternal.TypeA,
|
|
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
|
|
})
|
|
|
|
var events []TraversalEvent
|
|
tr.SetHooks(&TraverserHooks{
|
|
OnEvent: func(e TraversalEvent) {
|
|
events = append(events, e)
|
|
},
|
|
})
|
|
|
|
ctx := context.Background()
|
|
_, err := tr.Traverse(ctx, "example.com")
|
|
if err != nil {
|
|
t.Fatalf("Traverse: %v", err)
|
|
}
|
|
if len(events) == 0 {
|
|
t.Error("expected events from hooks")
|
|
}
|
|
}
|
|
|
|
func TestProcessReferralNoAddresses(t *testing.T) {
|
|
// Scenario: a referral without addresses. resolveGlueViaSystem fails (expired ctx),
|
|
// then ResolveNS is tried via the mock exchange.
|
|
answerMsg := new(dns.Msg)
|
|
answerMsg.SetReply(new(dns.Msg))
|
|
answerMsg.Answer = append(answerMsg.Answer, &dns.A{
|
|
Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300},
|
|
A: net.ParseIP("1.2.3.4"),
|
|
})
|
|
|
|
finalAnswerMsg := new(dns.Msg)
|
|
finalAnswerMsg.SetReply(new(dns.Msg))
|
|
finalAnswerMsg.Answer = append(finalAnswerMsg.Answer, &dns.A{
|
|
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300},
|
|
A: net.ParseIP("5.6.7.8"),
|
|
})
|
|
|
|
callCount := 0
|
|
tr := NewTraverser(&TraverserConfig{
|
|
MaxDepth: 5,
|
|
QueryType: dnsinternal.TypeA,
|
|
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++
|
|
q := msg.Question[0]
|
|
if q.Qtype == dns.TypeA && q.Name == "ns1.example.com." {
|
|
return answerMsg.Copy(), nil
|
|
}
|
|
return finalAnswerMsg.Copy(), nil
|
|
})
|
|
|
|
// Create a referral with no addresses (the NS name needs to be resolved)
|
|
ref := NewReferral("example.com.", dnsinternal.TypeA, "ns1.example.com.", 1, 1.0, nil)
|
|
// Do NOT set addresses - this exercises processReferral's no-address path
|
|
|
|
cache := NewInfoCache(nil)
|
|
// Use expired context for resolveGlueViaSystem so it returns nil fast
|
|
bgCtx := context.Background()
|
|
resp := tr.processReferral(bgCtx, ref, cache)
|
|
// Result may vary depending on whether 127.0.0.1:53 is available,
|
|
// but the function should not panic.
|
|
t.Logf("processReferral result type: %v", resp.Type)
|
|
}
|
|
|
|
func TestReferralResolveAlreadyHasAddresses(t *testing.T) {
|
|
ref := NewReferral("example.com.", dnsinternal.TypeA, ".", 0, 1.0, nil)
|
|
ref.Addresses = []net.IP{net.ParseIP("1.2.3.4")}
|
|
|
|
tr := NewTraverser(&TraverserConfig{
|
|
MaxDepth: 5,
|
|
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()
|
|
err := ref.Resolve(ctx, tr, nil, nil, 0)
|
|
if err != nil {
|
|
t.Fatalf("Resolve with existing addresses: %v", err)
|
|
}
|
|
if ref.State != StateResolved {
|
|
t.Errorf("expected StateResolved, got %v", ref.State)
|
|
}
|
|
}
|
|
|
|
func TestReferralResolveCacheHit(t *testing.T) {
|
|
ref := NewReferral("ns1.example.com.", dnsinternal.TypeA, ".", 0, 1.0, nil)
|
|
|
|
tr := NewTraverser(&TraverserConfig{
|
|
MaxDepth: 5,
|
|
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
|
|
})
|
|
|
|
cache := NewInfoCache(nil)
|
|
cache.StoreGlue("ns1.example.com.", []net.IP{net.ParseIP("1.2.3.4")})
|
|
|
|
ctx := context.Background()
|
|
err := ref.Resolve(ctx, tr, cache, nil, 0)
|
|
if err != nil {
|
|
t.Fatalf("Resolve cache hit: %v", err)
|
|
}
|
|
if ref.State != StateResolved {
|
|
t.Errorf("expected StateResolved, got %v", ref.State)
|
|
}
|
|
}
|
|
|
|
func TestReferralResolveCircular(t *testing.T) {
|
|
ref := NewReferral("ns1.example.com.", dnsinternal.TypeA, ".", 0, 1.0, nil)
|
|
|
|
tr := NewTraverser(&TraverserConfig{
|
|
MaxDepth: 5,
|
|
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
|
|
})
|
|
|
|
visited := map[string]bool{"ns1.example.com.": true}
|
|
ctx := context.Background()
|
|
err := ref.Resolve(ctx, tr, nil, visited, 0)
|
|
if err == nil {
|
|
t.Fatal("expected circular referral error")
|
|
}
|
|
if _, ok := err.(*CircularReferralError); !ok {
|
|
t.Errorf("expected CircularReferralError, got %T: %v", err, err)
|
|
}
|
|
}
|
|
|
|
func TestReferralResolveMaxDepth(t *testing.T) {
|
|
ref := NewReferral("ns1.example.com.", dnsinternal.TypeA, ".", 0, 1.0, nil)
|
|
|
|
tr := NewTraverser(&TraverserConfig{
|
|
MaxDepth: 5,
|
|
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()
|
|
err := ref.Resolve(ctx, tr, nil, nil, DefaultMaxDepth+1)
|
|
if err == nil {
|
|
t.Fatal("expected max depth error")
|
|
}
|
|
if _, ok := err.(*UnresolvableNameserverError); !ok {
|
|
t.Errorf("expected UnresolvableNameserverError, got %T: %v", err, err)
|
|
}
|
|
}
|
|
|
|
func TestReferralResolveWithAnswer(t *testing.T) {
|
|
answerMsg := new(dns.Msg)
|
|
answerMsg.SetReply(new(dns.Msg))
|
|
answerMsg.Answer = append(answerMsg.Answer, &dns.A{
|
|
Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300},
|
|
A: net.ParseIP("1.2.3.4"),
|
|
})
|
|
|
|
tr := NewTraverser(&TraverserConfig{
|
|
MaxDepth: 5,
|
|
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 answerMsg.Copy(), nil
|
|
})
|
|
|
|
ref := NewReferral("ns1.example.com.", dnsinternal.TypeA, ".", 0, 1.0, nil)
|
|
ctx := context.Background()
|
|
err := ref.Resolve(ctx, tr, nil, nil, 0)
|
|
if err != nil {
|
|
t.Fatalf("Resolve: %v", err)
|
|
}
|
|
if ref.State != StateResolved {
|
|
t.Errorf("expected StateResolved, got %v", ref.State)
|
|
}
|
|
if len(ref.Addresses) == 0 {
|
|
t.Error("expected addresses after resolution")
|
|
}
|
|
}
|
|
|
|
func TestReferralResolveNXDOMAIN(t *testing.T) {
|
|
nxMsg := new(dns.Msg)
|
|
nxMsg.Rcode = dns.RcodeNameError
|
|
|
|
tr := NewTraverser(&TraverserConfig{
|
|
MaxDepth: 5,
|
|
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 nxMsg.Copy(), nil
|
|
})
|
|
|
|
ref := NewReferral("nonexistent.invalid.", dnsinternal.TypeA, ".", 0, 1.0, nil)
|
|
ctx := context.Background()
|
|
err := ref.Resolve(ctx, tr, nil, nil, 0)
|
|
if err == nil {
|
|
t.Fatal("expected error for NXDOMAIN")
|
|
}
|
|
if _, ok := err.(*UnresolvableNameserverError); !ok {
|
|
t.Errorf("expected UnresolvableNameserverError, got %T: %v", err, err)
|
|
}
|
|
}
|
|
|
|
func TestReferralResolveContextCancellation(t *testing.T) {
|
|
tr := NewTraverser(&TraverserConfig{
|
|
MaxDepth: 5,
|
|
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) {
|
|
refMsg := new(dns.Msg)
|
|
refMsg.Rcode = dns.RcodeSuccess
|
|
refMsg.Ns = append(refMsg.Ns, &dns.NS{
|
|
Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS},
|
|
Ns: "ns.example.com.",
|
|
})
|
|
refMsg.Extra = append(refMsg.Extra, &dns.A{
|
|
Hdr: dns.RR_Header{Name: "ns.example.com.", Rrtype: dns.TypeA},
|
|
A: net.ParseIP("1.2.3.4"),
|
|
})
|
|
return refMsg, nil
|
|
})
|
|
|
|
ref := NewReferral("ns1.example.com.", dnsinternal.TypeA, ".", 0, 1.0, nil)
|
|
ctx, cancel := context.WithCancel(context.Background())
|
|
cancel() // Cancel immediately
|
|
|
|
err := ref.Resolve(ctx, tr, nil, nil, 0)
|
|
if err == nil {
|
|
t.Fatal("expected error on cancelled context")
|
|
}
|
|
}
|
|
|
|
func TestReferralResolveReferralPath(t *testing.T) {
|
|
// Test Resolve when it gets a referral response that pushes to stack
|
|
referralMsg := new(dns.Msg)
|
|
referralMsg.Rcode = dns.RcodeSuccess
|
|
referralMsg.Ns = append(referralMsg.Ns, &dns.NS{
|
|
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeNS, Class: dns.ClassINET},
|
|
Ns: "ns1.example.com.",
|
|
})
|
|
referralMsg.Extra = append(referralMsg.Extra, &dns.A{
|
|
Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET},
|
|
A: net.ParseIP("1.2.3.4"),
|
|
})
|
|
|
|
answerMsg := new(dns.Msg)
|
|
answerMsg.SetReply(new(dns.Msg))
|
|
answerMsg.Answer = append(answerMsg.Answer, &dns.A{
|
|
Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300},
|
|
A: net.ParseIP("5.6.7.8"),
|
|
})
|
|
|
|
callCount := 0
|
|
tr := NewTraverser(&TraverserConfig{
|
|
MaxDepth: 5,
|
|
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++
|
|
if callCount <= 1 {
|
|
return referralMsg.Copy(), nil
|
|
}
|
|
return answerMsg.Copy(), nil
|
|
})
|
|
|
|
ref := NewReferral("ns1.example.com.", dnsinternal.TypeA, ".", 0, 1.0, nil)
|
|
ctx := context.Background()
|
|
err := ref.Resolve(ctx, tr, nil, nil, 0)
|
|
// May succeed or exhaust depending on referral loop
|
|
t.Logf("Resolve referral path: err=%v, state=%v", err, ref.State)
|
|
}
|
|
|
|
func TestResolutionStateStringUnknown(t *testing.T) {
|
|
// Cover the default case of ResolutionState.String()
|
|
unknown := ResolutionState(99)
|
|
s := unknown.String()
|
|
if s != "unknown" {
|
|
t.Errorf("expected 'unknown' for invalid ResolutionState, got %q", s)
|
|
}
|
|
}
|
|
|
|
func TestTraverserNonFastMode(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: dns.TypeA, Class: dns.ClassINET, Ttl: 300},
|
|
A: net.ParseIP("1.2.3.4"),
|
|
})
|
|
|
|
tr := NewTraverser(&TraverserConfig{
|
|
MaxDepth: 5,
|
|
QueryType: dnsinternal.TypeA,
|
|
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
|
Fast: false, // Non-fast mode
|
|
})
|
|
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("Traverse non-fast: %v", err)
|
|
}
|
|
if len(results) == 0 {
|
|
t.Fatal("expected results")
|
|
}
|
|
}
|
|
|
|
func TestIterativeQueryWithExchangeUsesConfig(t *testing.T) {
|
|
answerMsg := new(dns.Msg)
|
|
answerMsg.SetReply(new(dns.Msg))
|
|
answerMsg.Answer = append(answerMsg.Answer, &dns.A{
|
|
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300},
|
|
A: net.ParseIP("1.2.3.4"),
|
|
})
|
|
|
|
tr := NewTraverser(&TraverserConfig{
|
|
MaxDepth: 5,
|
|
QueryType: dnsinternal.TypeA,
|
|
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
|
QueryConfig: dnsinternal.DefaultQueryConfig(),
|
|
})
|
|
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
|
return answerMsg.Copy(), nil
|
|
})
|
|
|
|
ctx := context.Background()
|
|
msg, err := tr.iterativeQueryWithExchange(ctx, net.ParseIP("198.41.0.4"), "example.com.", dnsinternal.TypeA)
|
|
if err != nil {
|
|
t.Fatalf("iterativeQueryWithExchange with config: %v", err)
|
|
}
|
|
if msg == nil {
|
|
t.Fatal("expected non-nil response")
|
|
}
|
|
}
|
|
|
|
func TestTraverserReferralWithHooks(t *testing.T) {
|
|
// Tests that hooks are called with IsResolve=true during ResolveNS sub-traversal
|
|
answerMsg := new(dns.Msg)
|
|
answerMsg.SetReply(new(dns.Msg))
|
|
answerMsg.Answer = append(answerMsg.Answer, &dns.A{
|
|
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300},
|
|
A: net.ParseIP("1.2.3.4"),
|
|
})
|
|
|
|
tr := NewTraverser(&TraverserConfig{
|
|
MaxDepth: 5,
|
|
QueryType: dnsinternal.TypeA,
|
|
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 answerMsg.Copy(), nil
|
|
})
|
|
|
|
var resolveEvents, progressEvents int
|
|
tr.SetHooks(&TraverserHooks{
|
|
OnEvent: func(e TraversalEvent) {
|
|
if e.IsResolve {
|
|
resolveEvents++
|
|
} else {
|
|
progressEvents++
|
|
}
|
|
},
|
|
})
|
|
|
|
// Test directly via ResolveNS with hooks
|
|
ctx := context.Background()
|
|
addrs, err := tr.ResolveNS(ctx, "ns1.example.com.", nil, nil, 0)
|
|
if err != nil {
|
|
t.Fatalf("ResolveNS: %v", err)
|
|
}
|
|
_ = addrs
|
|
t.Logf("resolveEvents=%d progressEvents=%d", resolveEvents, progressEvents)
|
|
}
|