feat: implement traversal engine (HAN-380) #6

Merged
multica-agent merged 1 commits from agent/go-expert-developer/9f18abb8 into main 2026-06-06 03:41:15 +00:00
11 changed files with 2098 additions and 0 deletions
Showing only changes of commit 2835ee86cf - Show all commits
+68
View File
@@ -125,6 +125,74 @@ func QueryWithExchange(ctx context.Context, server net.IP, name string, qtype ui
return nil, fmt.Errorf("query %s %s failed after %d retries: %w", name, QNameType(qtype), cfg.Retries, lastErr)
}
func IterativeQuery(ctx context.Context, server net.IP, name string, qtype uint16, cfg *QueryConfig) (*dns.Msg, error) {
if cfg == nil {
cfg = DefaultQueryConfig()
}
if cfg.UDPSize <= 0 {
cfg.UDPSize = DefaultEDNS0UDPSize()
}
return IterativeQueryWithExchange(ctx, server, name, qtype, cfg, realExchange)
}
func IterativeQueryWithExchange(ctx context.Context, server net.IP, name string, qtype uint16, cfg *QueryConfig, exchangeFn ExchangeFunc) (*dns.Msg, error) {
if cfg == nil {
cfg = DefaultQueryConfig()
}
if cfg.UDPSize <= 0 {
cfg.UDPSize = DefaultEDNS0UDPSize()
}
msg := buildQuery(name, qtype, cfg.UDPSize)
msg.RecursionDesired = false
serverStr := server.String()
var lastErr error
for attempt := 0; attempt < cfg.Retries; attempt++ {
if attempt > 0 {
select {
case <-ctx.Done():
return nil, fmt.Errorf("query retries cancelled: %w", ctx.Err())
case <-time.After(100 * time.Millisecond):
}
}
if cfg.UseTCP {
resp, err := exchangeFn(ctx, serverStr, msg, true)
if err != nil {
lastErr = err
continue
}
return resp, nil
}
resp, err := exchangeFn(ctx, serverStr, msg, false)
if err != nil {
lastErr = err
continue
}
if resp == nil {
lastErr = fmt.Errorf("nil response")
continue
}
if resp.Truncated {
resp, err = exchangeFn(ctx, serverStr, msg, true)
if err != nil {
lastErr = err
continue
}
return resp, nil
}
return resp, nil
}
return nil, fmt.Errorf("iterative query %s %s failed after %d retries: %w", name, QNameType(qtype), cfg.Retries, lastErr)
}
func buildQuery(name string, qtype uint16, udpSize int) *dns.Msg {
m := new(dns.Msg)
m.SetQuestion(dns.Fqdn(name), qtype)
+124
View File
@@ -0,0 +1,124 @@
package traverse
import (
"net"
"strings"
"sync"
miekgdns "github.com/miekg/dns"
)
type InfoCache struct {
parent *InfoCache
mu sync.RWMutex
ns map[string][]string
glue map[string][]net.IP
}
func NewInfoCache(parent *InfoCache) *InfoCache {
return &InfoCache{
parent: parent,
ns: make(map[string][]string),
glue: make(map[string][]net.IP),
}
}
func (c *InfoCache) StoreNS(zone string, nameservers []string) {
if len(nameservers) == 0 {
return
}
zone = normalize(zone)
c.mu.Lock()
seen := make(map[string]bool)
for _, ns := range nameservers {
ns = normalize(ns)
if !seen[ns] {
seen[ns] = true
c.ns[zone] = append(c.ns[zone], ns)
}
}
c.mu.Unlock()
}
func (c *InfoCache) LookupNS(zone string) []string {
zone = normalize(zone)
if names := c.localNS(zone); len(names) > 0 {
return names
}
if c.parent != nil {
return c.parent.LookupNS(zone)
}
return nil
}
func (c *InfoCache) localNS(zone string) []string {
c.mu.RLock()
defer c.mu.RUnlock()
names, ok := c.ns[zone]
if !ok {
return nil
}
result := make([]string, len(names))
copy(result, names)
return result
}
func (c *InfoCache) StoreGlue(name string, addrs []net.IP) {
if len(addrs) == 0 {
return
}
name = normalize(name)
c.mu.Lock()
seen := make(map[string]bool)
for _, addr := range addrs {
key := addr.String()
if !seen[key] {
seen[key] = true
c.glue[name] = append(c.glue[name], addr)
}
}
c.mu.Unlock()
}
func (c *InfoCache) LookupGlue(name string) []net.IP {
name = normalize(name)
if addrs := c.localGlue(name); len(addrs) > 0 {
return addrs
}
if c.parent != nil {
return c.parent.LookupGlue(name)
}
return nil
}
func (c *InfoCache) localGlue(name string) []net.IP {
c.mu.RLock()
defer c.mu.RUnlock()
addrs, ok := c.glue[name]
if !ok {
return nil
}
result := make([]net.IP, len(addrs))
copy(result, addrs)
return result
}
func (c *InfoCache) Child() *InfoCache {
return NewInfoCache(c)
}
func (c *InfoCache) NSCount() int {
c.mu.RLock()
defer c.mu.RUnlock()
return len(c.ns)
}
func (c *InfoCache) GlueCount() int {
c.mu.RLock()
defer c.mu.RUnlock()
return len(c.glue)
}
func normalize(name string) string {
return strings.ToLower(miekgdns.Fqdn(name))
}
+203
View File
@@ -0,0 +1,203 @@
package traverse
import (
"fmt"
"net"
"strings"
"sync"
"testing"
"github.com/miekg/dns"
)
func TestNewInfoCache(t *testing.T) {
c := NewInfoCache(nil)
if c.parent != nil {
t.Error("root cache should have nil parent")
}
if c.NSCount() != 0 {
t.Errorf("NSCount = %d, want 0", c.NSCount())
}
if c.GlueCount() != 0 {
t.Errorf("GlueCount = %d, want 0", c.GlueCount())
}
}
func TestInfoCacheStoreAndLookupNS(t *testing.T) {
c := NewInfoCache(nil)
c.StoreNS("com.", []string{"a.gtld-servers.net.", "b.gtld-servers.net."})
if c.NSCount() != 1 {
t.Errorf("NSCount = %d, want 1", c.NSCount())
}
names := c.LookupNS("com.")
if len(names) != 2 {
t.Fatalf("expected 2 nameservers, got %d", len(names))
}
if names[0] != "a.gtld-servers.net." {
t.Errorf("nameserver[0] = %q, want %q", names[0], "a.gtld-servers.net.")
}
}
func TestInfoCacheNSDedup(t *testing.T) {
c := NewInfoCache(nil)
c.StoreNS("com.", []string{"a.gtld-servers.net.", "a.gtld-servers.net."})
names := c.LookupNS("com.")
if len(names) != 1 {
t.Errorf("expected 1 deduped NS, got %d", len(names))
}
}
func TestInfoCacheNSCaseInsensitive(t *testing.T) {
c := NewInfoCache(nil)
c.StoreNS("COM.", []string{"A.GTLD-SERVERS.NET."})
names := c.LookupNS("com.")
if len(names) != 1 {
t.Fatalf("expected 1 NS, got %d", len(names))
}
if names[0] != "a.gtld-servers.net." {
t.Errorf("NS = %q, want %q", names[0], "a.gtld-servers.net.")
}
}
func TestInfoCacheNSLookupMiss(t *testing.T) {
c := NewInfoCache(nil)
names := c.LookupNS("org.")
if names != nil {
t.Errorf("expected nil for miss, got %v", names)
}
}
func TestInfoCacheNSStoreEmpty(t *testing.T) {
c := NewInfoCache(nil)
c.StoreNS("com.", nil)
if c.NSCount() != 0 {
t.Errorf("expected 0 after empty store, got %d", c.NSCount())
}
}
func TestInfoCacheChainedNS(t *testing.T) {
parent := NewInfoCache(nil)
parent.StoreNS("com.", []string{"a.gtld-servers.net."})
child := parent.Child()
if child.parent != parent {
t.Error("child parent should be the parent cache")
}
names := child.LookupNS("com.")
if len(names) != 1 {
t.Fatalf("expected 1 NS from parent, got %d", len(names))
}
if child.NSCount() != 0 {
t.Errorf("child NSCount = %d, want 0", child.NSCount())
}
}
func TestInfoCacheChildOverridesParent(t *testing.T) {
parent := NewInfoCache(nil)
parent.StoreNS("com.", []string{"a.gtld-servers.net."})
child := parent.Child()
child.StoreNS("com.", []string{"b.gtld-servers.net."})
names := child.LookupNS("com.")
if len(names) != 1 {
t.Fatalf("expected 1 NS, got %d", len(names))
}
if names[0] != "b.gtld-servers.net." {
t.Errorf("expected child's NS to override, got %q", names[0])
}
}
func TestInfoCacheStoreAndLookupGlue(t *testing.T) {
c := NewInfoCache(nil)
addrs := []net.IP{net.ParseIP("1.2.3.4"), net.ParseIP("5.6.7.8")}
c.StoreGlue("ns1.example.com.", addrs)
result := c.LookupGlue("ns1.example.com.")
if len(result) != 2 {
t.Fatalf("expected 2 glue addresses, got %d", len(result))
}
}
func TestInfoCacheGlueDedup(t *testing.T) {
c := NewInfoCache(nil)
ip := net.ParseIP("1.2.3.4")
c.StoreGlue("ns1.example.com.", []net.IP{ip, ip})
result := c.LookupGlue("ns1.example.com.")
if len(result) != 1 {
t.Errorf("expected 1 deduped glue, got %d", len(result))
}
}
func TestInfoCacheChainedGlue(t *testing.T) {
parent := NewInfoCache(nil)
parent.StoreGlue("ns1.example.com.", []net.IP{net.ParseIP("1.2.3.4")})
child := parent.Child()
result := child.LookupGlue("ns1.example.com.")
if len(result) != 1 {
t.Fatalf("expected 1 glue from parent, got %d", len(result))
}
if child.GlueCount() != 0 {
t.Errorf("child GlueCount = %d, want 0", child.GlueCount())
}
}
func TestInfoCacheGlueLookupMiss(t *testing.T) {
c := NewInfoCache(nil)
result := c.LookupGlue("nonexistent.example.com.")
if result != nil {
t.Errorf("expected nil for miss, got %v", result)
}
}
func TestInfoCacheConcurrentAccess(t *testing.T) {
c := NewInfoCache(nil)
var wg sync.WaitGroup
for i := 0; i < 100; i++ {
wg.Add(1)
go func(i int) {
defer wg.Done()
name := strings.ToLower(dns.Fqdn(fmt.Sprintf("ns%d.example.com.", i)))
c.StoreNS("example.com.", []string{name})
c.StoreGlue(name, []net.IP{net.ParseIP(fmt.Sprintf("1.2.3.%d", i%256))})
_ = c.LookupNS("example.com.")
_ = c.LookupGlue(name)
}(i)
}
wg.Wait()
}
func TestInfoCacheNilParent(t *testing.T) {
c := NewInfoCache(nil)
if c.LookupNS("com.") != nil {
t.Error("root cache should return nil for miss")
}
if c.LookupGlue("ns.example.com.") != nil {
t.Error("root cache should return nil for glue miss")
}
}
func TestNormalize(t *testing.T) {
tests := []struct {
input string
want string
}{
{"example.com", "example.com."},
{"Example.COM.", "example.com."},
{"EXAMPLE.COM", "example.com."},
}
for _, tt := range tests {
t.Run(tt.input, func(t *testing.T) {
got := normalize(tt.input)
if got != tt.want {
t.Errorf("normalize(%q) = %q, want %q", tt.input, got, tt.want)
}
})
}
}
+78
View File
@@ -0,0 +1,78 @@
package traverse
import (
"net"
"strings"
miekgdns "github.com/miekg/dns"
)
type ResolutionState int
const (
StateUnresolved ResolutionState = iota
StateResolving
StateResolved
)
func (s ResolutionState) String() string {
switch s {
case StateUnresolved:
return "unresolved"
case StateResolving:
return "resolving"
case StateResolved:
return "resolved"
default:
return "unknown"
}
}
type Referral struct {
Name string
Qtype uint16
Qclass uint16
Bailiwick string
Addresses []net.IP
State ResolutionState
NSName string
Parent *Referral
Depth int
Prob float64
}
func NewReferral(name string, qtype uint16, bailiwick string, depth int, prob float64, parent *Referral) *Referral {
return &Referral{
Name: miekgdns.Fqdn(strings.ToLower(name)),
Qtype: qtype,
Qclass: miekgdns.ClassINET,
Bailiwick: miekgdns.Fqdn(strings.ToLower(bailiwick)),
Depth: depth,
Prob: prob,
Parent: parent,
State: StateUnresolved,
}
}
func (r *Referral) InBailiwick(name string) bool {
if r.Bailiwick == "" || r.Bailiwick == "." {
return true
}
fqdn := miekgdns.Fqdn(strings.ToLower(name))
return miekgdns.IsSubDomain(r.Bailiwick, fqdn)
}
func (r *Referral) HasAddresses() bool {
return len(r.Addresses) > 0
}
func (r *Referral) SetAddresses(addrs []net.IP) {
r.Addresses = addrs
if len(addrs) > 0 {
r.State = StateResolved
} else {
r.State = StateUnresolved
}
}
+111
View File
@@ -0,0 +1,111 @@
package traverse
import (
"net"
"testing"
"github.com/hits/ExploreDNS/internal/dns"
)
const TypeA = dns.TypeA
type State = ResolutionState
func TestNewReferral(t *testing.T) {
ref := NewReferral("example.com", TypeA, ".", 0, 1.0, nil)
if ref.Name != "example.com." {
t.Errorf("Name = %q, want %q", ref.Name, "example.com.")
}
if ref.Qtype != TypeA {
t.Errorf("Qtype = %d, want %d", ref.Qtype, TypeA)
}
if ref.Qclass != 1 {
t.Errorf("Qclass = %d, want 1", ref.Qclass)
}
if ref.State != StateUnresolved {
t.Errorf("State = %d, want %d", ref.State, StateUnresolved)
}
if ref.Depth != 0 {
t.Errorf("Depth = %d, want 0", ref.Depth)
}
if ref.Prob != 1.0 {
t.Errorf("Prob = %f, want 1.0", ref.Prob)
}
if ref.Parent != nil {
t.Error("Parent should be nil")
}
}
func TestReferralInBailiwick(t *testing.T) {
tests := []struct {
name string
bailiwick string
testName string
want bool
}{
{"root bailiwick accepts all", ".", "example.com.", true},
{"empty bailiwick accepts all", "", "example.com.", true},
{"subdomain in bailiwick", "com.", "example.com.", true},
{"deeper subdomain", "com.", "www.example.com.", true},
{"not in bailiwick", "org.", "example.com.", false},
{"same zone", "example.com.", "example.com.", true},
{"sibling zone", "example.com.", "other.com.", false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
ref := &Referral{Bailiwick: tt.bailiwick}
if got := ref.InBailiwick(tt.testName); got != tt.want {
t.Errorf("InBailiwick(%q) = %v, want %v", tt.testName, got, tt.want)
}
})
}
}
func TestReferralHasAddresses(t *testing.T) {
ref := &Referral{}
if ref.HasAddresses() {
t.Error("empty referral should not have addresses")
}
ref.Addresses = []net.IP{net.ParseIP("1.2.3.4")}
if !ref.HasAddresses() {
t.Error("referral with address should have addresses")
}
}
func TestReferralSetAddresses(t *testing.T) {
ref := &Referral{}
ref.SetAddresses([]net.IP{net.ParseIP("1.2.3.4")})
if ref.State != StateResolved {
t.Errorf("State = %d, want %d", ref.State, StateResolved)
}
if !ref.HasAddresses() {
t.Error("should have addresses after SetAddresses")
}
ref.SetAddresses(nil)
if ref.State != StateUnresolved {
t.Errorf("State = %d, want %d", ref.State, StateUnresolved)
}
}
func TestResolutionStateString(t *testing.T) {
tests := []struct {
state State
want string
}{
{StateUnresolved, "unresolved"},
{StateResolving, "resolving"},
{StateResolved, "resolved"},
}
for _, tt := range tests {
t.Run(tt.want, func(t *testing.T) {
if got := tt.state.String(); got != tt.want {
t.Errorf("String() = %q, want %q", got, tt.want)
}
})
}
}
+218
View File
@@ -0,0 +1,218 @@
package traverse
import (
"net"
"github.com/hits/ExploreDNS/internal/dns"
miekgdns "github.com/miekg/dns"
)
type ResponseType int
const (
RespReferral ResponseType = iota
RespAnswer
RespCNAMEFollow
RespNODATA
RespNXDOMAIN
RespSERVFAIL
RespError
)
func (rt ResponseType) String() string {
switch rt {
case RespReferral:
return "referral"
case RespAnswer:
return "answer"
case RespCNAMEFollow:
return "cname_follow"
case RespNODATA:
return "nodata"
case RespNXDOMAIN:
return "nxdomain"
case RespSERVFAIL:
return "servfail"
case RespError:
return "error"
default:
return "unknown"
}
}
type Response struct {
Referral *Referral
Server net.IP
Cache *InfoCache
Decoded *dns.DecodedResponse
Type ResponseType
}
func NewResponse(ref *Referral, server net.IP, cache *InfoCache) *Response {
return &Response{
Referral: ref,
Server: server,
Cache: cache,
}
}
func (r *Response) Process(msg *miekgdns.Msg) *Response {
if msg == nil {
r.Type = RespError
return r
}
r.Decoded = dns.DecodeResponse(msg)
if r.Decoded == nil {
r.Type = RespError
return r
}
r.Type = r.classify()
return r
}
func (r *Response) classify() ResponseType {
switch r.Decoded.Classification {
case dns.ResponseNXDOMAIN:
return RespNXDOMAIN
case dns.ResponseSERVFAIL:
return RespSERVFAIL
case dns.ResponseAnswer:
if len(r.Decoded.CNAMEChain) > 0 && !r.hasFinalAnswer() {
return RespCNAMEFollow
}
return RespAnswer
case dns.ResponseReferral:
return RespReferral
case dns.ResponseNODATA:
return RespNODATA
default:
return RespError
}
}
func (r *Response) hasFinalAnswer() bool {
for _, rr := range r.Decoded.Answers {
if _, ok := rr.(*miekgdns.CNAME); ok {
continue
}
return true
}
return false
}
func (r *Response) ChildReferrals() []*Referral {
if r.Type != RespReferral {
return nil
}
if r.Referral == nil {
return nil
}
var nameservers []string
for _, rr := range r.Decoded.Authority {
if ns, ok := rr.(*miekgdns.NS); ok {
if r.Referral.InBailiwick(ns.Ns) {
nameservers = append(nameservers, ns.Ns)
}
}
}
if len(nameservers) == 0 {
for _, rr := range r.Decoded.Authority {
if ns, ok := rr.(*miekgdns.NS); ok {
nameservers = append(nameservers, ns.Ns)
}
}
}
r.storeAuthority(nameservers)
prob := r.childProb(len(nameservers))
var children []*Referral
for _, ns := range nameservers {
child := NewReferral(
r.Referral.Name,
r.Referral.Qtype,
ns,
r.Referral.Depth+1,
prob,
r.Referral,
)
r.resolveGlue(child)
children = append(children, child)
}
return children
}
func (r *Response) CNAMEFollowReferral() *Referral {
if r.Type != RespCNAMEFollow || len(r.Decoded.CNAMEChain) == 0 {
return nil
}
target := r.Decoded.CNAMEChain[len(r.Decoded.CNAMEChain)-1]
follow := NewReferral(
target,
r.Referral.Qtype,
r.Referral.Bailiwick,
r.Referral.Depth+1,
r.Referral.Prob,
r.Referral,
)
if len(r.Referral.Addresses) > 0 {
follow.Addresses = make([]net.IP, len(r.Referral.Addresses))
copy(follow.Addresses, r.Referral.Addresses)
follow.State = StateResolved
}
return follow
}
func (r *Response) storeAuthority(nameservers []string) {
if r.Cache == nil {
return
}
zone := r.Referral.Name
r.Cache.StoreNS(zone, nameservers)
}
func (r *Response) resolveGlue(child *Referral) {
if r.Cache == nil {
return
}
nsName := child.Bailiwick
for _, rr := range r.Decoded.Additional {
switch v := rr.(type) {
case *miekgdns.A:
if normalize(v.Header().Name) == normalize(nsName) {
child.Addresses = append(child.Addresses, v.A)
}
case *miekgdns.AAAA:
if normalize(v.Header().Name) == normalize(nsName) {
child.Addresses = append(child.Addresses, v.AAAA)
}
}
}
if child.HasAddresses() {
child.State = StateResolved
}
r.Cache.StoreGlue(nsName, child.Addresses)
}
func (r *Response) IsTerminal() bool {
switch r.Type {
case RespAnswer, RespNODATA, RespNXDOMAIN, RespSERVFAIL, RespError:
return true
default:
return false
}
}
func (r *Response) childProb(n int) float64 {
if n <= 0 {
return 0
}
if r.Referral == nil {
return 1.0 / float64(n)
}
return r.Referral.Prob / float64(n)
}
+338
View File
@@ -0,0 +1,338 @@
package traverse
import (
"net"
"testing"
"github.com/miekg/dns"
)
func TestResponseProcessNil(t *testing.T) {
ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil)
r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil)
r.Process(nil)
if r.Type != RespError {
t.Errorf("Type = %d, want %d", r.Type, RespError)
}
}
func TestResponseClassifyAnswer(t *testing.T) {
msg := new(dns.Msg)
msg.SetReply(new(dns.Msg))
msg.Answer = append(msg.Answer, &dns.A{
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300},
A: net.ParseIP("93.184.216.34"),
})
ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil)
r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil)
r.Process(msg)
if r.Type != RespAnswer {
t.Errorf("Type = %d, want %d", r.Type, RespAnswer)
}
if r.Decoded == nil {
t.Fatal("Decoded should not be nil")
}
}
func TestResponseClassifyReferral(t *testing.T) {
msg := new(dns.Msg)
msg.Rcode = dns.RcodeSuccess
msg.Authoritative = false
msg.Ns = append(msg.Ns, &dns.NS{
Hdr: dns.RR_Header{Name: "com.", Rrtype: dns.TypeNS, Class: dns.ClassINET, Ttl: 172800},
Ns: "a.gtld-servers.net.",
})
msg.Extra = append(msg.Extra, &dns.A{
Hdr: dns.RR_Header{Name: "a.gtld-servers.net.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 172800},
A: net.ParseIP("192.5.6.30"),
})
ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil)
cache := NewInfoCache(nil)
r := NewResponse(ref, net.ParseIP("1.2.3.4"), cache)
r.Process(msg)
if r.Type != RespReferral {
t.Errorf("Type = %d, want %d", r.Type, RespReferral)
}
}
func TestResponseClassifyNXDOMAIN(t *testing.T) {
msg := new(dns.Msg)
msg.Rcode = dns.RcodeNameError
ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil)
r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil)
r.Process(msg)
if r.Type != RespNXDOMAIN {
t.Errorf("Type = %d, want %d", r.Type, RespNXDOMAIN)
}
}
func TestResponseClassifyNODATA(t *testing.T) {
msg := new(dns.Msg)
msg.Rcode = dns.RcodeSuccess
msg.Authoritative = true
msg.Ns = append(msg.Ns, &dns.SOA{
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeSOA, Class: dns.ClassINET, Ttl: 3600},
})
ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil)
r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil)
r.Process(msg)
if r.Type != RespNODATA {
t.Errorf("Type = %d, want %d", r.Type, RespNODATA)
}
}
func TestResponseClassifySERVFAIL(t *testing.T) {
msg := new(dns.Msg)
msg.Rcode = dns.RcodeServerFailure
ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil)
r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil)
r.Process(msg)
if r.Type != RespSERVFAIL {
t.Errorf("Type = %d, want %d", r.Type, RespSERVFAIL)
}
}
func TestResponseCNAMEFollow(t *testing.T) {
msg := new(dns.Msg)
msg.SetReply(new(dns.Msg))
msg.Answer = append(msg.Answer,
&dns.CNAME{
Hdr: dns.RR_Header{Name: "www.example.com.", Rrtype: dns.TypeCNAME, Class: dns.ClassINET},
Target: "example.com.",
},
)
ref := NewReferral("www.example.com", dns.TypeA, ".", 0, 1.0, nil)
r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil)
r.Process(msg)
if r.Type != RespCNAMEFollow {
t.Errorf("Type = %d, want %d", r.Type, RespCNAMEFollow)
}
}
func TestResponseCNAMEWithFinalAnswer(t *testing.T) {
msg := new(dns.Msg)
msg.SetReply(new(dns.Msg))
msg.Answer = append(msg.Answer,
&dns.CNAME{
Hdr: dns.RR_Header{Name: "www.example.com.", Rrtype: dns.TypeCNAME, Class: dns.ClassINET},
Target: "example.com.",
},
&dns.A{
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300},
A: net.ParseIP("93.184.216.34"),
},
)
ref := NewReferral("www.example.com", dns.TypeA, ".", 0, 1.0, nil)
r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil)
r.Process(msg)
if r.Type != RespAnswer {
t.Errorf("Type = %d, want %d (CNAME with final A answer)", r.Type, RespAnswer)
}
}
func TestResponseChildReferrals(t *testing.T) {
msg := new(dns.Msg)
msg.Rcode = dns.RcodeSuccess
msg.Authoritative = false
msg.Ns = append(msg.Ns,
&dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeNS}, Ns: "a.gtld-servers.net."},
&dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeNS}, Ns: "b.gtld-servers.net."},
)
msg.Extra = append(msg.Extra,
&dns.A{Hdr: dns.RR_Header{Name: "a.gtld-servers.net.", Rrtype: dns.TypeA}, A: net.ParseIP("192.5.6.30")},
&dns.A{Hdr: dns.RR_Header{Name: "b.gtld-servers.net.", Rrtype: dns.TypeA}, A: net.ParseIP("192.33.14.30")},
)
ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil)
cache := NewInfoCache(nil)
r := NewResponse(ref, net.ParseIP("198.41.0.4"), cache)
r.Process(msg)
children := r.ChildReferrals()
if len(children) != 2 {
t.Fatalf("expected 2 child referrals, got %d", len(children))
}
if children[0].Name != "example.com." {
t.Errorf("child[0] name = %q, want example.com.", children[0].Name)
}
if children[0].Bailiwick != "a.gtld-servers.net." {
t.Errorf("child[0] bailiwick = %q, want a.gtld-servers.net.", children[0].Bailiwick)
}
if children[0].Prob != 0.5 {
t.Errorf("child[0] prob = %f, want 0.5", children[0].Prob)
}
if children[0].Depth != 1 {
t.Errorf("child[0] depth = %d, want 1", children[0].Depth)
}
if !children[0].HasAddresses() {
t.Error("child[0] should have glue addresses")
}
if !children[1].HasAddresses() {
t.Error("child[1] should have glue addresses")
}
nsNames := cache.LookupNS("example.com.")
if len(nsNames) != 2 {
t.Errorf("expected 2 NS in cache, got %d", len(nsNames))
}
}
func TestResponseChildReferralsNonReferral(t *testing.T) {
msg := new(dns.Msg)
msg.SetReply(new(dns.Msg))
msg.Answer = append(msg.Answer, &dns.A{
Hdr: dns.RR_Header{Rrtype: dns.TypeA},
A: net.ParseIP("1.2.3.4"),
})
ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil)
r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil)
r.Process(msg)
if children := r.ChildReferrals(); children != nil {
t.Error("non-referral should not produce child referrals")
}
}
func TestResponseCNAMEFollowReferral(t *testing.T) {
msg := new(dns.Msg)
msg.SetReply(new(dns.Msg))
msg.Answer = append(msg.Answer,
&dns.CNAME{
Hdr: dns.RR_Header{Name: "www.example.com.", Rrtype: dns.TypeCNAME},
Target: "example.com.",
},
)
ref := NewReferral("www.example.com", dns.TypeA, ".", 0, 1.0, nil)
r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil)
r.Process(msg)
follow := r.CNAMEFollowReferral()
if follow == nil {
t.Fatal("expected CNAME follow referral")
}
if follow.Name != "example.com." {
t.Errorf("follow name = %q, want %q", follow.Name, "example.com.")
}
if follow.Depth != 1 {
t.Errorf("follow depth = %d, want 1", follow.Depth)
}
}
func TestResponseIsTerminal(t *testing.T) {
tests := []struct {
respType ResponseType
want bool
}{
{RespAnswer, true},
{RespNODATA, true},
{RespNXDOMAIN, true},
{RespSERVFAIL, true},
{RespError, true},
{RespReferral, false},
{RespCNAMEFollow, false},
}
for _, tt := range tests {
t.Run(tt.respType.String(), func(t *testing.T) {
r := &Response{Type: tt.respType}
if got := r.IsTerminal(); got != tt.want {
t.Errorf("IsTerminal() = %v, want %v", got, tt.want)
}
})
}
}
func TestResponseTypeString(t *testing.T) {
tests := []struct {
rt ResponseType
want string
}{
{RespReferral, "referral"},
{RespAnswer, "answer"},
{RespCNAMEFollow, "cname_follow"},
{RespNODATA, "nodata"},
{RespNXDOMAIN, "nxdomain"},
{RespSERVFAIL, "servfail"},
{RespError, "error"},
}
for _, tt := range tests {
t.Run(tt.want, func(t *testing.T) {
if got := tt.rt.String(); got != tt.want {
t.Errorf("String() = %q, want %q", got, tt.want)
}
})
}
}
func TestResponseChildReferralsProbabilityInheritance(t *testing.T) {
msg := new(dns.Msg)
msg.Rcode = dns.RcodeSuccess
msg.Authoritative = false
msg.Ns = append(msg.Ns,
&dns.NS{Hdr: dns.RR_Header{Name: "com.", Rrtype: dns.TypeNS}, Ns: "a.gtld-servers.net."},
&dns.NS{Hdr: dns.RR_Header{Name: "com.", Rrtype: dns.TypeNS}, Ns: "b.gtld-servers.net."},
&dns.NS{Hdr: dns.RR_Header{Name: "com.", Rrtype: dns.TypeNS}, Ns: "c.gtld-servers.net."},
)
ref := NewReferral("example.com", dns.TypeA, ".", 0, 0.5, nil)
r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil)
r.Process(msg)
children := r.ChildReferrals()
if len(children) != 3 {
t.Fatalf("expected 3 children, got %d", len(children))
}
for _, c := range children {
if c.Prob != 0.5/3.0 {
t.Errorf("child prob = %f, want %f", c.Prob, 0.5/3.0)
}
}
}
func TestResponseChildReferralsEmptyAuthority(t *testing.T) {
msg := new(dns.Msg)
msg.Rcode = dns.RcodeSuccess
msg.Authoritative = false
ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil)
r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil)
r.Process(msg)
children := r.ChildReferrals()
if len(children) != 0 {
t.Errorf("expected 0 children with empty authority, got %d", len(children))
}
}
func TestResponseNilReferral(t *testing.T) {
r := NewResponse(nil, net.ParseIP("1.2.3.4"), nil)
children := r.ChildReferrals()
if children != nil {
t.Error("nil referral should produce no children")
}
follow := r.CNAMEFollowReferral()
if follow != nil {
t.Error("nil referral should produce no CNAME follow")
}
}
+58
View File
@@ -0,0 +1,58 @@
package traverse
const DefaultMaxDepth = 20
type Stack struct {
items []*Referral
maxDepth int
}
func NewStack(maxDepth int) *Stack {
if maxDepth <= 0 {
maxDepth = DefaultMaxDepth
}
return &Stack{
items: make([]*Referral, 0),
maxDepth: maxDepth,
}
}
func (s *Stack) Push(r *Referral) bool {
if r == nil {
return false
}
if r.Depth >= s.maxDepth {
return false
}
s.items = append(s.items, r)
return true
}
func (s *Stack) Pop() *Referral {
if len(s.items) == 0 {
return nil
}
idx := len(s.items) - 1
item := s.items[idx]
s.items = s.items[:idx]
return item
}
func (s *Stack) Peek() *Referral {
if len(s.items) == 0 {
return nil
}
return s.items[len(s.items)-1]
}
func (s *Stack) Len() int {
return len(s.items)
}
func (s *Stack) MaxDepth() int {
return s.maxDepth
}
func (s *Stack) IsEmpty() bool {
return len(s.items) == 0
}
+143
View File
@@ -0,0 +1,143 @@
package traverse
import (
"testing"
"github.com/hits/ExploreDNS/internal/dns"
)
func TestNewStack(t *testing.T) {
s := NewStack(10)
if s.MaxDepth() != 10 {
t.Errorf("MaxDepth = %d, want 10", s.MaxDepth())
}
if !s.IsEmpty() {
t.Error("new stack should be empty")
}
if s.Len() != 0 {
t.Errorf("Len = %d, want 0", s.Len())
}
}
func TestNewStackDefaultDepth(t *testing.T) {
s := NewStack(0)
if s.MaxDepth() != DefaultMaxDepth {
t.Errorf("MaxDepth = %d, want %d", s.MaxDepth(), DefaultMaxDepth)
}
s = NewStack(-5)
if s.MaxDepth() != DefaultMaxDepth {
t.Errorf("MaxDepth = %d, want %d", s.MaxDepth(), DefaultMaxDepth)
}
}
func TestStackPushPop(t *testing.T) {
s := NewStack(5)
ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil)
ok := s.Push(ref)
if !ok {
t.Error("Push should succeed")
}
if s.Len() != 1 {
t.Errorf("Len = %d, want 1", s.Len())
}
popped := s.Pop()
if popped != ref {
t.Error("popped referral should match pushed")
}
if s.Len() != 0 {
t.Errorf("Len = %d, want 0", s.Len())
}
}
func TestStackLIFO(t *testing.T) {
s := NewStack(5)
r1 := NewReferral("a.com", dns.TypeA, ".", 0, 1.0, nil)
r2 := NewReferral("b.com", dns.TypeA, ".", 1, 1.0, nil)
r3 := NewReferral("c.com", dns.TypeA, ".", 2, 1.0, nil)
s.Push(r1)
s.Push(r2)
s.Push(r3)
if popped := s.Pop(); popped != r3 {
t.Error("should pop r3 first (LIFO)")
}
if popped := s.Pop(); popped != r2 {
t.Error("should pop r2 second")
}
if popped := s.Pop(); popped != r1 {
t.Error("should pop r1 third")
}
}
func TestStackPushNil(t *testing.T) {
s := NewStack(5)
ok := s.Push(nil)
if ok {
t.Error("Push(nil) should return false")
}
if s.Len() != 0 {
t.Errorf("Len = %d, want 0", s.Len())
}
}
func TestStackMaxDepth(t *testing.T) {
s := NewStack(3)
r0 := NewReferral("a.com", dns.TypeA, ".", 0, 1.0, nil)
r1 := NewReferral("b.com", dns.TypeA, ".", 1, 1.0, nil)
r2 := NewReferral("c.com", dns.TypeA, ".", 2, 1.0, nil)
r3 := NewReferral("d.com", dns.TypeA, ".", 3, 1.0, nil)
if !s.Push(r0) {
t.Error("depth 0 should be accepted")
}
if !s.Push(r1) {
t.Error("depth 1 should be accepted")
}
if !s.Push(r2) {
t.Error("depth 2 should be accepted")
}
if s.Push(r3) {
t.Error("depth 3 should be rejected (maxDepth=3)")
}
}
func TestStackPopEmpty(t *testing.T) {
s := NewStack(5)
if popped := s.Pop(); popped != nil {
t.Error("Pop on empty stack should return nil")
}
}
func TestStackPeek(t *testing.T) {
s := NewStack(5)
if peek := s.Peek(); peek != nil {
t.Error("Peek on empty stack should return nil")
}
ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil)
s.Push(ref)
if peek := s.Peek(); peek != ref {
t.Error("Peek should return top item")
}
if s.Len() != 1 {
t.Errorf("Peek should not remove item, Len = %d, want 1", s.Len())
}
}
func TestStackIsEmpty(t *testing.T) {
s := NewStack(5)
if !s.IsEmpty() {
t.Error("new stack should be empty")
}
s.Push(NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil))
if s.IsEmpty() {
t.Error("stack with item should not be empty")
}
}
+274
View File
@@ -0,0 +1,274 @@
package traverse
import (
"context"
"fmt"
"net"
"sync"
"github.com/hits/ExploreDNS/internal/dns"
miekgdns "github.com/miekg/dns"
)
type TraverserConfig struct {
MaxDepth int
QueryType uint16
RootConfig *dns.RootDiscoveryConfig
QueryConfig *dns.QueryConfig
RootAddrs []net.IP
}
func DefaultTraverserConfig() *TraverserConfig {
return &TraverserConfig{
MaxDepth: DefaultMaxDepth,
QueryType: dns.TypeA,
RootConfig: nil,
QueryConfig: nil,
RootAddrs: nil,
}
}
type TraversalResult struct {
Referral *Referral
Response *Response
}
type Traverser struct {
config *TraverserConfig
exchange dns.ExchangeFunc
}
func NewTraverser(cfg *TraverserConfig) *Traverser {
if cfg == nil {
cfg = DefaultTraverserConfig()
}
return &Traverser{
config: cfg,
exchange: nil,
}
}
func (t *Traverser) SetExchange(fn dns.ExchangeFunc) {
t.exchange = fn
}
func (t *Traverser) Traverse(ctx context.Context, name string) ([]TraversalResult, error) {
name = miekgdns.Fqdn(name)
roots, err := t.discoverRoots(ctx)
if err != nil {
return nil, fmt.Errorf("root discovery: %w", err)
}
initial := NewReferral(name, t.config.QueryType, ".", 0, 1.0, nil)
initial.Addresses = roots
stack := NewStack(t.config.MaxDepth)
stack.Push(initial)
rootCache := NewInfoCache(nil)
var (
mu sync.Mutex
results []TraversalResult
)
for {
select {
case <-ctx.Done():
return results, fmt.Errorf("traversal cancelled: %w", ctx.Err())
default:
}
ref := stack.Pop()
if ref == nil {
break
}
cache := rootCache
if ref.Parent != nil {
cache = rootCache.Child()
}
resp := t.processReferral(ctx, ref, cache)
mu.Lock()
results = append(results, TraversalResult{Referral: ref, Response: resp})
mu.Unlock()
if resp.IsTerminal() {
continue
}
if resp.Type == RespReferral {
children := resp.ChildReferrals()
for _, child := range children {
if !stack.Push(child) {
mu.Lock()
results = append(results, TraversalResult{
Referral: child,
Response: &Response{
Referral: child,
Type: RespError,
},
})
mu.Unlock()
}
}
}
if resp.Type == RespCNAMEFollow {
follow := resp.CNAMEFollowReferral()
if follow != nil {
if !stack.Push(follow) {
mu.Lock()
results = append(results, TraversalResult{
Referral: follow,
Response: &Response{
Referral: follow,
Type: RespError,
},
})
mu.Unlock()
}
}
}
}
return results, nil
}
func (t *Traverser) discoverRoots(ctx context.Context) ([]net.IP, error) {
if len(t.config.RootAddrs) > 0 {
return t.config.RootAddrs, nil
}
servers, err := dns.DiscoverRoots(ctx, t.config.RootConfig)
if err != nil {
return nil, err
}
var addrs []net.IP
for _, srv := range servers {
addrs = append(addrs, srv.AllIPs(false)...)
}
return addrs, nil
}
func (t *Traverser) processReferral(ctx context.Context, ref *Referral, cache *InfoCache) *Response {
if !ref.HasAddresses() {
ref.Addresses = t.resolveGlueViaSystem(ctx, ref.Name, cache)
if len(ref.Addresses) > 0 {
ref.State = StateResolved
} else {
return &Response{
Referral: ref,
Type: RespError,
}
}
}
for _, addr := range ref.Addresses {
resp := t.queryServer(ctx, ref, addr, cache)
if resp != nil && resp.Type != RespSERVFAIL {
return resp
}
}
return &Response{
Referral: ref,
Type: RespSERVFAIL,
}
}
func (t *Traverser) queryServer(ctx context.Context, ref *Referral, server net.IP, cache *InfoCache) *Response {
var msg *miekgdns.Msg
var err error
if t.exchange != nil {
msg, err = t.iterativeQueryWithExchange(ctx, server, ref.Name, ref.Qtype)
} else {
msg, err = dns.Query(ctx, server, ref.Name, ref.Qtype, t.config.QueryConfig)
if err == nil {
msg = t.ensureRDFalse(msg, server, ref.Name, ref.Qtype, t.config.QueryConfig)
}
}
if err != nil {
return &Response{
Referral: ref,
Server: server,
Type: RespError,
}
}
resp := NewResponse(ref, server, cache)
resp.Process(msg)
return resp
}
func (t *Traverser) iterativeQueryWithExchange(ctx context.Context, server net.IP, name string, qtype uint16) (*miekgdns.Msg, error) {
if t.config.QueryConfig == nil {
return dns.IterativeQueryWithExchange(ctx, server, name, qtype, nil, t.exchange)
}
return dns.IterativeQueryWithExchange(ctx, server, name, qtype, t.config.QueryConfig, t.exchange)
}
func (t *Traverser) ensureRDFalse(msg *miekgdns.Msg, server net.IP, name string, qtype uint16, cfg *dns.QueryConfig) *miekgdns.Msg {
if msg != nil && msg.RecursionDesired {
if t.exchange != nil {
ctx := context.Background()
var err error
msg, err = t.iterativeQueryWithExchange(ctx, server, name, qtype)
if err != nil {
return nil
}
return msg
}
msg.RecursionDesired = false
}
return msg
}
func (t *Traverser) resolveGlueViaSystem(ctx context.Context, name string, cache *InfoCache) []net.IP {
if cache != nil {
if addrs := cache.LookupGlue(name); len(addrs) > 0 {
return addrs
}
}
c := &miekgdns.Client{
Net: "udp",
ReadTimeout: 5,
WriteTimeout: 5,
}
if deadline, ok := ctx.Deadline(); ok {
c.ReadTimeout = deadline.Sub(deadline)
c.WriteTimeout = deadline.Sub(deadline)
}
fqdn := miekgdns.Fqdn(name)
aMsg, _, err := c.ExchangeContext(ctx, newAQuery(fqdn), "127.0.0.1:53")
if err == nil {
var addrs []net.IP
for _, rr := range aMsg.Answer {
if a, ok := rr.(*miekgdns.A); ok {
addrs = append(addrs, a.A)
}
}
if len(addrs) > 0 {
if cache != nil {
cache.StoreGlue(name, addrs)
}
return addrs
}
}
return nil
}
func newAQuery(name string) *miekgdns.Msg {
m := new(miekgdns.Msg)
m.SetQuestion(name, miekgdns.TypeA)
m.RecursionDesired = true
return m
}
+483
View File
@@ -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")
}
}