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