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:
@@ -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)
|
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 {
|
func buildQuery(name string, qtype uint16, udpSize int) *dns.Msg {
|
||||||
m := new(dns.Msg)
|
m := new(dns.Msg)
|
||||||
m.SetQuestion(dns.Fqdn(name), qtype)
|
m.SetQuestion(dns.Fqdn(name), qtype)
|
||||||
|
|||||||
@@ -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