feat: rework engine and CLI for dnstraverse parity
Port the traversal engine to the Ruby dnstraverse model so behaviour and output match dns.squish.net: - dns: single RD=0 query path (RD=1 only for upstream root discovery), per-run packet cache, EDNS0 512-fallback with warnings, UDP->TCP on truncation; fix --retries 0 and --root-server IP-literal handling; drop all hardcoded 127.0.0.1:53 resolvers - traverse: hierarchical per-branch InfoCache, 7-step response classification with the full 10-status vocabulary, bailiwick partitioning, strictly-deeper lame-referral rule, refid grammar with .0 resolve subtrees and childset digits, per-IP branching at 1/n weight, cache-based glue resolution with noglue/loop dead ends, CNAME restarts from the deepest cached zone, fast-mode memoization, probability aggregation with Ruby-identical stats keys (sums to 1.0) - output: byte-for-byte reference text format pinned by a golden test, reference CLI defaults, working --quiet/--show-X=false, TTY-aware colour, deduplicated deterministic JSON - web: adapt API/SPA to the new engine, SSE events carry refid/status, fix subscribe/snapshot duplicate-event race and a statusCls TDZ bug, align SPA type list with the backend - delete the old engine and dead code (net -4,350 lines) Verified against live runs of the reference Ruby engine across five domains (answers, NXDOMAIN, null MX, CNAME restart, glueless resolve) with no divergences beyond the documented typo fixes. Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Fable 5
parent
af15c9c2d4
commit
d71c7fbef2
+135
-94
@@ -1,6 +1,7 @@
|
||||
package traverse
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net"
|
||||
"strings"
|
||||
"sync"
|
||||
@@ -8,117 +9,157 @@ import (
|
||||
miekgdns "github.com/miekg/dns"
|
||||
)
|
||||
|
||||
// StartServer is one entry returned by GetStartServers: a nameserver hostname
|
||||
// plus its cached IPv4 addresses. IPs == nil means no addresses are cached and
|
||||
// the caller must resolve the name itself (glueless).
|
||||
type StartServer struct {
|
||||
Name string
|
||||
IPs []string
|
||||
}
|
||||
|
||||
// InfoCache is the hierarchical per-branch record cache (info_cache.rb). Each
|
||||
// response wraps its parent's cache in a child so sibling branches never see
|
||||
// each other's records; lookups recurse towards the root cache.
|
||||
type InfoCache struct {
|
||||
parent *InfoCache
|
||||
mu sync.RWMutex
|
||||
ns map[string][]string
|
||||
glue map[string][]net.IP
|
||||
data map[string][]miekgdns.RR
|
||||
}
|
||||
|
||||
func NewInfoCache(parent *InfoCache) *InfoCache {
|
||||
return &InfoCache{
|
||||
parent: parent,
|
||||
ns: make(map[string][]string),
|
||||
glue: make(map[string][]net.IP),
|
||||
data: make(map[string][]miekgdns.RR),
|
||||
}
|
||||
}
|
||||
|
||||
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)
|
||||
// canonicalName lowercases a DNS name and strips the trailing dot; the root
|
||||
// (and empty string) canonicalises to "", matching the Ruby engine's
|
||||
// representation of "no bailiwick".
|
||||
func canonicalName(name string) string {
|
||||
return strings.TrimSuffix(strings.ToLower(name), ".")
|
||||
}
|
||||
|
||||
func (c *InfoCache) GlueCount() int {
|
||||
c.mu.RLock()
|
||||
defer c.mu.RUnlock()
|
||||
return len(c.glue)
|
||||
func cacheKey(name string, qclass, qtype uint16) string {
|
||||
return fmt.Sprintf("%s:%d:%d", canonicalName(name), qclass, qtype)
|
||||
}
|
||||
|
||||
func normalize(name string) string {
|
||||
return strings.ToLower(miekgdns.Fqdn(name))
|
||||
func rrCacheKey(rr miekgdns.RR) string {
|
||||
h := rr.Header()
|
||||
return cacheKey(h.Name, h.Class, h.Rrtype)
|
||||
}
|
||||
|
||||
// Add stores resource records, REPLACING any existing entries that share a
|
||||
// name:class:type key (info_cache.rb add: clear pass, then append pass, so
|
||||
// several records under one key in a single call are all kept).
|
||||
func (c *InfoCache) Add(rrs []miekgdns.RR) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
for _, rr := range rrs {
|
||||
c.data[rrCacheKey(rr)] = nil
|
||||
}
|
||||
for _, rr := range rrs {
|
||||
key := rrCacheKey(rr)
|
||||
c.data[key] = append(c.data[key], rr)
|
||||
}
|
||||
}
|
||||
|
||||
// AddHints seeds NS records for domain ("" = root hints) plus A/AAAA records
|
||||
// for each server that has known addresses (info_cache.rb add_hints).
|
||||
func (c *InfoCache) AddHints(domain string, servers []StartServer) {
|
||||
var rrs []miekgdns.RR
|
||||
owner := miekgdns.Fqdn(canonicalName(domain))
|
||||
for _, srv := range servers {
|
||||
name := miekgdns.Fqdn(canonicalName(srv.Name))
|
||||
rrs = append(rrs, &miekgdns.NS{
|
||||
Hdr: miekgdns.RR_Header{Name: owner, Rrtype: miekgdns.TypeNS, Class: miekgdns.ClassINET},
|
||||
Ns: name,
|
||||
})
|
||||
for _, ip := range srv.IPs {
|
||||
addr := net.ParseIP(ip)
|
||||
if addr == nil {
|
||||
continue
|
||||
}
|
||||
if v4 := addr.To4(); v4 != nil {
|
||||
rrs = append(rrs, &miekgdns.A{
|
||||
Hdr: miekgdns.RR_Header{Name: name, Rrtype: miekgdns.TypeA, Class: miekgdns.ClassINET},
|
||||
A: v4,
|
||||
})
|
||||
} else {
|
||||
rrs = append(rrs, &miekgdns.AAAA{
|
||||
Hdr: miekgdns.RR_Header{Name: name, Rrtype: miekgdns.TypeAAAA, Class: miekgdns.ClassINET},
|
||||
AAAA: addr,
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
c.Add(rrs)
|
||||
}
|
||||
|
||||
// Get returns the cached RRset for name/class/type, consulting parent caches
|
||||
// on a local miss. Returns nil when nothing is cached anywhere in the chain.
|
||||
func (c *InfoCache) Get(name string, qclass, qtype uint16) []miekgdns.RR {
|
||||
key := cacheKey(name, qclass, qtype)
|
||||
c.mu.RLock()
|
||||
rrs, ok := c.data[key]
|
||||
c.mu.RUnlock()
|
||||
if ok {
|
||||
out := make([]miekgdns.RR, len(rrs))
|
||||
copy(out, rrs)
|
||||
return out
|
||||
}
|
||||
if c.parent != nil {
|
||||
return c.parent.Get(name, qclass, qtype)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// getNS finds the nearest cached NS RRset at or above domain, walking labels
|
||||
// upward to the root (info_cache.rb get_ns?).
|
||||
func (c *InfoCache) getNS(domain string) ([]miekgdns.RR, error) {
|
||||
domain = canonicalName(domain)
|
||||
for {
|
||||
if rrs := c.Get(domain, miekgdns.ClassINET, miekgdns.TypeNS); len(rrs) > 0 {
|
||||
return rrs, nil
|
||||
}
|
||||
if domain == "" {
|
||||
return nil, fmt.Errorf("no nameservers available for %q -- no root hints set??", domain)
|
||||
}
|
||||
if i := strings.Index(domain, "."); i >= 0 {
|
||||
domain = domain[i+1:]
|
||||
} else {
|
||||
domain = ""
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// GetStartServers returns the servers to start querying for domain: the
|
||||
// nearest cached NS RRset walking labels upward, each nameserver paired with
|
||||
// its cached A addresses (nil when unknown). newbailiwick is the owner name of
|
||||
// that NS RRset ("" for root).
|
||||
func (c *InfoCache) GetStartServers(domain string) (starters []StartServer, newbailiwick string, err error) {
|
||||
ns, err := c.getNS(domain)
|
||||
if err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
for _, rr := range ns {
|
||||
nsrr, ok := rr.(*miekgdns.NS)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
name := canonicalName(nsrr.Ns)
|
||||
var ips []string
|
||||
for _, iprr := range c.Get(name, miekgdns.ClassINET, miekgdns.TypeA) {
|
||||
if a, ok := iprr.(*miekgdns.A); ok {
|
||||
ips = append(ips, a.A.String())
|
||||
}
|
||||
}
|
||||
starters = append(starters, StartServer{Name: name, IPs: ips})
|
||||
}
|
||||
newbailiwick = canonicalName(ns[0].Header().Name)
|
||||
return starters, newbailiwick, nil
|
||||
}
|
||||
|
||||
+188
-131
@@ -3,201 +3,258 @@ 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 nsRR(zone, target string) dns.RR {
|
||||
return &dns.NS{
|
||||
Hdr: dns.RR_Header{Name: dns.Fqdn(zone), Rrtype: dns.TypeNS, Class: dns.ClassINET},
|
||||
Ns: dns.Fqdn(target),
|
||||
}
|
||||
}
|
||||
|
||||
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 aRR(name, ip string) dns.RR {
|
||||
return &dns.A{
|
||||
Hdr: dns.RR_Header{Name: dns.Fqdn(name), Rrtype: dns.TypeA, Class: dns.ClassINET},
|
||||
A: net.ParseIP(ip).To4(),
|
||||
}
|
||||
}
|
||||
|
||||
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 aaaaRR(name, ip string) dns.RR {
|
||||
return &dns.AAAA{
|
||||
Hdr: dns.RR_Header{Name: dns.Fqdn(name), Rrtype: dns.TypeAAAA, Class: dns.ClassINET},
|
||||
AAAA: net.ParseIP(ip),
|
||||
}
|
||||
}
|
||||
|
||||
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))
|
||||
func TestCanonicalName(t *testing.T) {
|
||||
tests := []struct{ in, want string }{
|
||||
{"example.com", "example.com"},
|
||||
{"Example.COM.", "example.com"},
|
||||
{".", ""},
|
||||
{"", ""},
|
||||
{"WWW.Example.Com", "www.example.com"},
|
||||
}
|
||||
if names[0] != "a.gtld-servers.net." {
|
||||
t.Errorf("NS = %q, want %q", names[0], "a.gtld-servers.net.")
|
||||
for _, tt := range tests {
|
||||
if got := canonicalName(tt.in); got != tt.want {
|
||||
t.Errorf("canonicalName(%q) = %q, want %q", tt.in, got, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestInfoCacheNSLookupMiss(t *testing.T) {
|
||||
func TestInfoCacheAddAndGet(t *testing.T) {
|
||||
c := NewInfoCache(nil)
|
||||
names := c.LookupNS("org.")
|
||||
if names != nil {
|
||||
t.Errorf("expected nil for miss, got %v", names)
|
||||
c.Add([]dns.RR{nsRR("com", "a.gtld-servers.net"), nsRR("com", "b.gtld-servers.net")})
|
||||
|
||||
rrs := c.Get("com", dns.ClassINET, dns.TypeNS)
|
||||
if len(rrs) != 2 {
|
||||
t.Fatalf("expected 2 NS records, got %d", len(rrs))
|
||||
}
|
||||
}
|
||||
|
||||
func TestInfoCacheNSStoreEmpty(t *testing.T) {
|
||||
func TestInfoCacheAddReplacesSameKey(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())
|
||||
c.Add([]dns.RR{nsRR("com", "old1.example.net"), nsRR("com", "old2.example.net")})
|
||||
c.Add([]dns.RR{nsRR("com", "new.example.net")})
|
||||
|
||||
rrs := c.Get("com", dns.ClassINET, dns.TypeNS)
|
||||
if len(rrs) != 1 {
|
||||
t.Fatalf("add should replace same name:class:type entry, got %d records", len(rrs))
|
||||
}
|
||||
if rrs[0].(*dns.NS).Ns != "new.example.net." {
|
||||
t.Errorf("NS = %q, want new.example.net.", rrs[0].(*dns.NS).Ns)
|
||||
}
|
||||
}
|
||||
|
||||
func TestInfoCacheChainedNS(t *testing.T) {
|
||||
func TestInfoCacheAddKeepsDistinctKeys(t *testing.T) {
|
||||
c := NewInfoCache(nil)
|
||||
c.Add([]dns.RR{nsRR("com", "a.gtld-servers.net"), aRR("a.gtld-servers.net", "192.5.6.30")})
|
||||
c.Add([]dns.RR{nsRR("org", "a0.org-servers.net")})
|
||||
|
||||
if got := c.Get("com", dns.ClassINET, dns.TypeNS); len(got) != 1 {
|
||||
t.Errorf("com NS lost after unrelated add: %v", got)
|
||||
}
|
||||
if got := c.Get("a.gtld-servers.net", dns.ClassINET, dns.TypeA); len(got) != 1 {
|
||||
t.Errorf("glue lost after unrelated add: %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestInfoCacheGetCaseInsensitive(t *testing.T) {
|
||||
c := NewInfoCache(nil)
|
||||
c.Add([]dns.RR{nsRR("COM.", "A.GTLD-SERVERS.NET.")})
|
||||
if got := c.Get("com", dns.ClassINET, dns.TypeNS); len(got) != 1 {
|
||||
t.Fatalf("expected case-insensitive hit, got %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestInfoCacheGetMiss(t *testing.T) {
|
||||
c := NewInfoCache(nil)
|
||||
if got := c.Get("org", dns.ClassINET, dns.TypeNS); got != nil {
|
||||
t.Errorf("expected nil for miss, got %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestInfoCacheGetRecursesToParent(t *testing.T) {
|
||||
parent := NewInfoCache(nil)
|
||||
parent.StoreNS("com.", []string{"a.gtld-servers.net."})
|
||||
|
||||
parent.Add([]dns.RR{nsRR("com", "a.gtld-servers.net")})
|
||||
child := parent.Child()
|
||||
if child.parent != parent {
|
||||
t.Error("child parent should be the parent cache")
|
||||
t.Fatal("Child() should link to parent")
|
||||
}
|
||||
|
||||
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())
|
||||
if got := child.Get("com", dns.ClassINET, dns.TypeNS); len(got) != 1 {
|
||||
t.Fatalf("expected parent hit through child, got %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestInfoCacheChildOverridesParent(t *testing.T) {
|
||||
func TestInfoCacheChildShadowsParent(t *testing.T) {
|
||||
parent := NewInfoCache(nil)
|
||||
parent.StoreNS("com.", []string{"a.gtld-servers.net."})
|
||||
|
||||
parent.Add([]dns.RR{nsRR("com", "parent.example.net")})
|
||||
child := parent.Child()
|
||||
child.StoreNS("com.", []string{"b.gtld-servers.net."})
|
||||
child.Add([]dns.RR{nsRR("com", "child.example.net")})
|
||||
|
||||
names := child.LookupNS("com.")
|
||||
if len(names) != 1 {
|
||||
t.Fatalf("expected 1 NS, got %d", len(names))
|
||||
got := child.Get("com", dns.ClassINET, dns.TypeNS)
|
||||
if len(got) != 1 || got[0].(*dns.NS).Ns != "child.example.net." {
|
||||
t.Errorf("child entry should shadow parent, got %v", got)
|
||||
}
|
||||
if names[0] != "b.gtld-servers.net." {
|
||||
t.Errorf("expected child's NS to override, got %q", names[0])
|
||||
// The parent must be untouched.
|
||||
got = parent.Get("com", dns.ClassINET, dns.TypeNS)
|
||||
if len(got) != 1 || got[0].(*dns.NS).Ns != "parent.example.net." {
|
||||
t.Errorf("parent entry modified, got %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestInfoCacheStoreAndLookupGlue(t *testing.T) {
|
||||
func TestGetStartServersWalksLabels(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)
|
||||
c.Add([]dns.RR{
|
||||
nsRR("com", "a.gtld-servers.net"),
|
||||
nsRR("com", "b.gtld-servers.net"),
|
||||
aRR("a.gtld-servers.net", "192.5.6.30"),
|
||||
})
|
||||
|
||||
result := c.LookupGlue("ns1.example.com.")
|
||||
if len(result) != 2 {
|
||||
t.Fatalf("expected 2 glue addresses, got %d", len(result))
|
||||
starters, bw, err := c.GetStartServers("www.deep.example.com")
|
||||
if err != nil {
|
||||
t.Fatalf("GetStartServers: %v", err)
|
||||
}
|
||||
if bw != "com" {
|
||||
t.Errorf("newbailiwick = %q, want com", bw)
|
||||
}
|
||||
if len(starters) != 2 {
|
||||
t.Fatalf("expected 2 starters, got %d", len(starters))
|
||||
}
|
||||
if starters[0].Name != "a.gtld-servers.net" {
|
||||
t.Errorf("starter[0] = %q", starters[0].Name)
|
||||
}
|
||||
if len(starters[0].IPs) != 1 || starters[0].IPs[0] != "192.5.6.30" {
|
||||
t.Errorf("starter[0] IPs = %v, want [192.5.6.30]", starters[0].IPs)
|
||||
}
|
||||
if starters[1].IPs != nil {
|
||||
t.Errorf("glueless starter should have nil IPs, got %v", starters[1].IPs)
|
||||
}
|
||||
}
|
||||
|
||||
func TestInfoCacheGlueDedup(t *testing.T) {
|
||||
func TestGetStartServersPrefersDeepestZone(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))
|
||||
c.AddHints("", []StartServer{{Name: "a.root-servers.net", IPs: []string{"198.41.0.4"}}})
|
||||
c.Add([]dns.RR{nsRR("example.com", "ns1.example.com"), aRR("ns1.example.com", "1.2.3.4")})
|
||||
|
||||
starters, bw, err := c.GetStartServers("www.example.com")
|
||||
if err != nil {
|
||||
t.Fatalf("GetStartServers: %v", err)
|
||||
}
|
||||
if bw != "example.com" {
|
||||
t.Errorf("newbailiwick = %q, want example.com", bw)
|
||||
}
|
||||
if len(starters) != 1 || starters[0].Name != "ns1.example.com" {
|
||||
t.Errorf("starters = %v", starters)
|
||||
}
|
||||
}
|
||||
|
||||
func TestInfoCacheChainedGlue(t *testing.T) {
|
||||
func TestGetStartServersRootHints(t *testing.T) {
|
||||
c := NewInfoCache(nil)
|
||||
c.AddHints("", []StartServer{
|
||||
{Name: "a.root-servers.net", IPs: []string{"198.41.0.4", "2001:503:ba3e::2:30"}},
|
||||
{Name: "b.root-servers.net", IPs: []string{"170.247.170.2"}},
|
||||
})
|
||||
|
||||
starters, bw, err := c.GetStartServers("anything.example.org")
|
||||
if err != nil {
|
||||
t.Fatalf("GetStartServers: %v", err)
|
||||
}
|
||||
if bw != "" {
|
||||
t.Errorf("root bailiwick should be \"\", got %q", bw)
|
||||
}
|
||||
if len(starters) != 2 {
|
||||
t.Fatalf("expected 2 root starters, got %d", len(starters))
|
||||
}
|
||||
// Only the IPv4 address surfaces (IPv4-only transport); the AAAA is
|
||||
// cached but not returned as a start address.
|
||||
if len(starters[0].IPs) != 1 || starters[0].IPs[0] != "198.41.0.4" {
|
||||
t.Errorf("starter[0].IPs = %v, want [198.41.0.4]", starters[0].IPs)
|
||||
}
|
||||
if got := c.Get("a.root-servers.net", dns.ClassINET, dns.TypeAAAA); len(got) != 1 {
|
||||
t.Errorf("AAAA hint should be cached, got %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetStartServersNoRootHints(t *testing.T) {
|
||||
c := NewInfoCache(nil)
|
||||
if _, _, err := c.GetStartServers("example.com"); err == nil {
|
||||
t.Fatal("expected error with no NS cached anywhere")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetStartServersExactDomainMatch(t *testing.T) {
|
||||
c := NewInfoCache(nil)
|
||||
c.Add([]dns.RR{nsRR("example.com", "ns1.example.net")})
|
||||
_, bw, err := c.GetStartServers("example.com")
|
||||
if err != nil {
|
||||
t.Fatalf("GetStartServers: %v", err)
|
||||
}
|
||||
if bw != "example.com" {
|
||||
t.Errorf("newbailiwick = %q, want example.com", bw)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetStartServersUsesBranchCache(t *testing.T) {
|
||||
parent := NewInfoCache(nil)
|
||||
parent.StoreGlue("ns1.example.com.", []net.IP{net.ParseIP("1.2.3.4")})
|
||||
|
||||
parent.AddHints("", []StartServer{{Name: "a.root-servers.net", IPs: []string{"198.41.0.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())
|
||||
}
|
||||
}
|
||||
child.Add([]dns.RR{nsRR("com", "a.gtld-servers.net")})
|
||||
|
||||
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)
|
||||
_, bw, err := child.GetStartServers("www.example.com")
|
||||
if err != nil {
|
||||
t.Fatalf("GetStartServers: %v", err)
|
||||
}
|
||||
if bw != "com" {
|
||||
t.Errorf("newbailiwick = %q, want com (child cache hit)", bw)
|
||||
}
|
||||
|
||||
// A sibling branch must not see the child's records.
|
||||
sibling := parent.Child()
|
||||
_, bw, err = sibling.GetStartServers("www.example.com")
|
||||
if err != nil {
|
||||
t.Fatalf("GetStartServers: %v", err)
|
||||
}
|
||||
if bw != "" {
|
||||
t.Errorf("sibling newbailiwick = %q, want \"\" (root only)", bw)
|
||||
}
|
||||
}
|
||||
|
||||
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)
|
||||
name := fmt.Sprintf("ns%d.example.com", i)
|
||||
c.Add([]dns.RR{nsRR("example.com", name), aRR(name, fmt.Sprintf("1.2.3.%d", i%256))})
|
||||
_, _, _ = c.GetStartServers("www.example.com")
|
||||
_ = c.Get(name, dns.ClassINET, dns.TypeA)
|
||||
}(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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,727 +0,0 @@
|
||||
package traverse
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
dnsinternal "gitea.hansenits.com.au/hits/ExploreDNS/internal/dns"
|
||||
"github.com/miekg/dns"
|
||||
)
|
||||
|
||||
func TestSetHooks(t *testing.T) {
|
||||
tr := NewTraverser(nil)
|
||||
hooks := &TraverserHooks{
|
||||
OnEvent: func(event TraversalEvent) {},
|
||||
}
|
||||
tr.SetHooks(hooks)
|
||||
if tr.config.Hooks != hooks {
|
||||
t.Error("SetHooks should set config.Hooks")
|
||||
}
|
||||
|
||||
// SetHooks on nil config traverser (initializes config)
|
||||
tr2 := &Traverser{}
|
||||
tr2.SetHooks(hooks)
|
||||
if tr2.config == nil || tr2.config.Hooks != hooks {
|
||||
t.Error("SetHooks should initialize config when nil")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewAQuery(t *testing.T) {
|
||||
msg := newAQuery("example.com.")
|
||||
if msg == nil {
|
||||
t.Fatal("newAQuery returned nil")
|
||||
}
|
||||
if !msg.RecursionDesired {
|
||||
t.Error("expected RD=true in newAQuery")
|
||||
}
|
||||
if len(msg.Question) == 0 {
|
||||
t.Fatal("expected question in newAQuery")
|
||||
}
|
||||
if msg.Question[0].Qtype != dns.TypeA {
|
||||
t.Errorf("expected TypeA, got %d", msg.Question[0].Qtype)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveGlueViaSystemCacheHit(t *testing.T) {
|
||||
tr := NewTraverser(&TraverserConfig{
|
||||
MaxDepth: 5,
|
||||
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||
})
|
||||
cache := NewInfoCache(nil)
|
||||
expected := []net.IP{net.ParseIP("1.2.3.4")}
|
||||
cache.StoreGlue("ns1.example.com.", expected)
|
||||
|
||||
ctx := context.Background()
|
||||
addrs := tr.resolveGlueViaSystem(ctx, "ns1.example.com.", cache)
|
||||
if len(addrs) == 0 {
|
||||
t.Error("expected addresses from cache hit")
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveGlueViaSystemExpiredContext(t *testing.T) {
|
||||
tr := NewTraverser(&TraverserConfig{
|
||||
MaxDepth: 5,
|
||||
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||
})
|
||||
|
||||
// Expired context → remaining <= 0 → returns nil immediately
|
||||
ctx, cancel := context.WithDeadline(context.Background(), time.Now().Add(-time.Second))
|
||||
defer cancel()
|
||||
|
||||
addrs := tr.resolveGlueViaSystem(ctx, "ns1.example.com.", nil)
|
||||
if len(addrs) != 0 {
|
||||
t.Errorf("expected nil from expired context, got %v", addrs)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveGlueViaSystemTimeout(t *testing.T) {
|
||||
tr := NewTraverser(&TraverserConfig{
|
||||
MaxDepth: 5,
|
||||
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||
})
|
||||
|
||||
// Very short timeout will fail the DNS query to 127.0.0.1:53
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Millisecond)
|
||||
defer cancel()
|
||||
time.Sleep(15 * time.Millisecond) // ensure it's expired
|
||||
|
||||
addrs := tr.resolveGlueViaSystem(ctx, "ns1.example.com.", nil)
|
||||
// May return nil (timeout) or addresses (if local resolver responds instantly)
|
||||
t.Logf("resolveGlueViaSystem returned %d addresses", len(addrs))
|
||||
}
|
||||
|
||||
func TestEnsureRDFalseWithExchange(t *testing.T) {
|
||||
rdFalseMsg := new(dns.Msg)
|
||||
rdFalseMsg.SetReply(new(dns.Msg))
|
||||
rdFalseMsg.RecursionDesired = false
|
||||
|
||||
tr := NewTraverser(&TraverserConfig{
|
||||
MaxDepth: 5,
|
||||
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||
})
|
||||
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||
return rdFalseMsg.Copy(), nil
|
||||
})
|
||||
|
||||
rdTrueMsg := new(dns.Msg)
|
||||
rdTrueMsg.SetReply(new(dns.Msg))
|
||||
rdTrueMsg.RecursionDesired = true
|
||||
|
||||
result := tr.ensureRDFalse(rdTrueMsg, net.ParseIP("198.41.0.4"), "example.com.", dnsinternal.TypeA, nil)
|
||||
if result == nil {
|
||||
t.Fatal("ensureRDFalse with exchange should return non-nil")
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnsureRDFalseNilMsg(t *testing.T) {
|
||||
tr := NewTraverser(&TraverserConfig{
|
||||
MaxDepth: 5,
|
||||
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||
})
|
||||
result := tr.ensureRDFalse(nil, net.ParseIP("1.2.3.4"), "example.com.", dnsinternal.TypeA, nil)
|
||||
if result != nil {
|
||||
t.Error("ensureRDFalse(nil) should return nil")
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnsureRDFalseRDAlreadyFalse(t *testing.T) {
|
||||
tr := NewTraverser(&TraverserConfig{
|
||||
MaxDepth: 5,
|
||||
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||
})
|
||||
|
||||
msg := new(dns.Msg)
|
||||
msg.RecursionDesired = false
|
||||
result := tr.ensureRDFalse(msg, net.ParseIP("1.2.3.4"), "example.com.", dnsinternal.TypeA, nil)
|
||||
if result != msg {
|
||||
t.Error("ensureRDFalse should return same msg when RD=false")
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnsureRDFalseNoExchange(t *testing.T) {
|
||||
tr := NewTraverser(&TraverserConfig{
|
||||
MaxDepth: 5,
|
||||
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||
})
|
||||
// No exchange set
|
||||
|
||||
msg := new(dns.Msg)
|
||||
msg.RecursionDesired = true
|
||||
result := tr.ensureRDFalse(msg, net.ParseIP("1.2.3.4"), "example.com.", dnsinternal.TypeA, nil)
|
||||
if result == nil {
|
||||
t.Fatal("ensureRDFalse without exchange should return msg with RD cleared")
|
||||
}
|
||||
if result.RecursionDesired {
|
||||
t.Error("expected RD=false after ensureRDFalse without exchange")
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveNSFromCache(t *testing.T) {
|
||||
tr := NewTraverser(&TraverserConfig{
|
||||
MaxDepth: 5,
|
||||
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||
})
|
||||
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||
return nil, nil
|
||||
})
|
||||
|
||||
cache := NewInfoCache(nil)
|
||||
expected := []net.IP{net.ParseIP("1.2.3.4")}
|
||||
cache.StoreGlue("ns1.example.com.", expected)
|
||||
|
||||
ctx := context.Background()
|
||||
addrs, err := tr.ResolveNS(ctx, "ns1.example.com.", cache, nil, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("ResolveNS cache hit: %v", err)
|
||||
}
|
||||
if len(addrs) == 0 {
|
||||
t.Error("expected addresses from cache")
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveNSCircularReferral(t *testing.T) {
|
||||
tr := NewTraverser(&TraverserConfig{
|
||||
MaxDepth: 5,
|
||||
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||
})
|
||||
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||
return nil, nil
|
||||
})
|
||||
|
||||
visited := map[string]bool{"ns1.example.com.": true}
|
||||
ctx := context.Background()
|
||||
_, err := tr.ResolveNS(ctx, "ns1.example.com.", nil, visited, 0)
|
||||
if err == nil {
|
||||
t.Fatal("expected circular referral error")
|
||||
}
|
||||
var circErr *CircularReferralError
|
||||
if _, ok := err.(*CircularReferralError); !ok {
|
||||
t.Errorf("expected CircularReferralError, got %T: %v", err, err)
|
||||
}
|
||||
_ = circErr
|
||||
}
|
||||
|
||||
func TestResolveNSMaxDepth(t *testing.T) {
|
||||
tr := NewTraverser(&TraverserConfig{
|
||||
MaxDepth: 5,
|
||||
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||
})
|
||||
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||
return nil, nil
|
||||
})
|
||||
|
||||
ctx := context.Background()
|
||||
_, err := tr.ResolveNS(ctx, "ns1.example.com.", nil, nil, DefaultMaxDepth+1)
|
||||
if err == nil {
|
||||
t.Fatal("expected max depth error")
|
||||
}
|
||||
if _, ok := err.(*UnresolvableNameserverError); !ok {
|
||||
t.Errorf("expected UnresolvableNameserverError, got %T: %v", err, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveNSWithAnswer(t *testing.T) {
|
||||
answerMsg := new(dns.Msg)
|
||||
answerMsg.SetReply(new(dns.Msg))
|
||||
answerMsg.Answer = append(answerMsg.Answer, &dns.A{
|
||||
Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300},
|
||||
A: net.ParseIP("1.2.3.4"),
|
||||
})
|
||||
|
||||
tr := NewTraverser(&TraverserConfig{
|
||||
MaxDepth: 5,
|
||||
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||
})
|
||||
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||
return answerMsg.Copy(), nil
|
||||
})
|
||||
|
||||
ctx := context.Background()
|
||||
addrs, err := tr.ResolveNS(ctx, "ns1.example.com.", nil, nil, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("ResolveNS with answer: %v", err)
|
||||
}
|
||||
if len(addrs) == 0 {
|
||||
t.Fatal("expected addresses from NS resolution")
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveNSNXDOMAIN(t *testing.T) {
|
||||
nxMsg := new(dns.Msg)
|
||||
nxMsg.Rcode = dns.RcodeNameError
|
||||
|
||||
tr := NewTraverser(&TraverserConfig{
|
||||
MaxDepth: 5,
|
||||
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||
})
|
||||
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||
return nxMsg.Copy(), nil
|
||||
})
|
||||
|
||||
ctx := context.Background()
|
||||
_, err := tr.ResolveNS(ctx, "nonexistent.invalid.", nil, nil, 0)
|
||||
if err == nil {
|
||||
t.Fatal("expected error for NXDOMAIN NS resolution")
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveNSContextCancellation(t *testing.T) {
|
||||
tr := NewTraverser(&TraverserConfig{
|
||||
MaxDepth: 5,
|
||||
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||
})
|
||||
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||
// Keep returning referrals to keep the loop going
|
||||
refMsg := new(dns.Msg)
|
||||
refMsg.Rcode = dns.RcodeSuccess
|
||||
refMsg.Ns = append(refMsg.Ns, &dns.NS{
|
||||
Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS, Class: dns.ClassINET},
|
||||
Ns: "ns.example.com.",
|
||||
})
|
||||
refMsg.Extra = append(refMsg.Extra, &dns.A{
|
||||
Hdr: dns.RR_Header{Name: "ns.example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET},
|
||||
A: net.ParseIP("1.2.3.4"),
|
||||
})
|
||||
return refMsg, nil
|
||||
})
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel() // Cancel immediately
|
||||
|
||||
_, err := tr.ResolveNS(ctx, "ns1.example.com.", nil, nil, 0)
|
||||
if err == nil {
|
||||
t.Fatal("expected error on cancelled context")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiscoverRootsWithRootAddrs(t *testing.T) {
|
||||
expected := []net.IP{net.ParseIP("198.41.0.4"), net.ParseIP("199.9.14.201")}
|
||||
tr := NewTraverser(&TraverserConfig{
|
||||
MaxDepth: 5,
|
||||
RootAddrs: expected,
|
||||
})
|
||||
|
||||
ctx := context.Background()
|
||||
addrs, err := tr.discoverRoots(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("discoverRoots with RootAddrs: %v", err)
|
||||
}
|
||||
if len(addrs) != len(expected) {
|
||||
t.Errorf("expected %d addresses, got %d", len(expected), len(addrs))
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiscoverRootsFromSystem(t *testing.T) {
|
||||
tr := NewTraverser(&TraverserConfig{
|
||||
MaxDepth: 5,
|
||||
// No RootAddrs - will call dns.DiscoverRoots
|
||||
})
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer cancel()
|
||||
|
||||
addrs, err := tr.discoverRoots(ctx)
|
||||
if err != nil {
|
||||
t.Logf("discoverRoots without RootAddrs error (may skip): %v", err)
|
||||
t.Skip()
|
||||
}
|
||||
if len(addrs) == 0 {
|
||||
t.Error("expected at least one root address")
|
||||
}
|
||||
}
|
||||
|
||||
func TestTraverserSetHooksAndTraverse(t *testing.T) {
|
||||
answerResp := new(dns.Msg)
|
||||
answerResp.SetReply(new(dns.Msg))
|
||||
answerResp.Answer = append(answerResp.Answer, &dns.A{
|
||||
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300},
|
||||
A: net.ParseIP("1.2.3.4"),
|
||||
})
|
||||
|
||||
tr := NewTraverser(&TraverserConfig{
|
||||
MaxDepth: 5,
|
||||
QueryType: dnsinternal.TypeA,
|
||||
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||
})
|
||||
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||
return answerResp.Copy(), nil
|
||||
})
|
||||
|
||||
var events []TraversalEvent
|
||||
tr.SetHooks(&TraverserHooks{
|
||||
OnEvent: func(e TraversalEvent) {
|
||||
events = append(events, e)
|
||||
},
|
||||
})
|
||||
|
||||
ctx := context.Background()
|
||||
_, err := tr.Traverse(ctx, "example.com")
|
||||
if err != nil {
|
||||
t.Fatalf("Traverse: %v", err)
|
||||
}
|
||||
if len(events) == 0 {
|
||||
t.Error("expected events from hooks")
|
||||
}
|
||||
}
|
||||
|
||||
func TestProcessReferralNoAddresses(t *testing.T) {
|
||||
// Scenario: a referral without addresses. resolveGlueViaSystem fails (expired ctx),
|
||||
// then ResolveNS is tried via the mock exchange.
|
||||
answerMsg := new(dns.Msg)
|
||||
answerMsg.SetReply(new(dns.Msg))
|
||||
answerMsg.Answer = append(answerMsg.Answer, &dns.A{
|
||||
Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300},
|
||||
A: net.ParseIP("1.2.3.4"),
|
||||
})
|
||||
|
||||
finalAnswerMsg := new(dns.Msg)
|
||||
finalAnswerMsg.SetReply(new(dns.Msg))
|
||||
finalAnswerMsg.Answer = append(finalAnswerMsg.Answer, &dns.A{
|
||||
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300},
|
||||
A: net.ParseIP("5.6.7.8"),
|
||||
})
|
||||
|
||||
callCount := 0
|
||||
tr := NewTraverser(&TraverserConfig{
|
||||
MaxDepth: 5,
|
||||
QueryType: dnsinternal.TypeA,
|
||||
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||
})
|
||||
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||
callCount++
|
||||
q := msg.Question[0]
|
||||
if q.Qtype == dns.TypeA && q.Name == "ns1.example.com." {
|
||||
return answerMsg.Copy(), nil
|
||||
}
|
||||
return finalAnswerMsg.Copy(), nil
|
||||
})
|
||||
|
||||
// Create a referral with no addresses (the NS name needs to be resolved)
|
||||
ref := NewReferral("example.com.", dnsinternal.TypeA, "ns1.example.com.", 1, 1.0, nil)
|
||||
// Do NOT set addresses - this exercises processReferral's no-address path
|
||||
|
||||
cache := NewInfoCache(nil)
|
||||
// Use expired context for resolveGlueViaSystem so it returns nil fast
|
||||
bgCtx := context.Background()
|
||||
resp := tr.processReferral(bgCtx, ref, cache)
|
||||
// Result may vary depending on whether 127.0.0.1:53 is available,
|
||||
// but the function should not panic.
|
||||
t.Logf("processReferral result type: %v", resp.Type)
|
||||
}
|
||||
|
||||
func TestReferralResolveAlreadyHasAddresses(t *testing.T) {
|
||||
ref := NewReferral("example.com.", dnsinternal.TypeA, ".", 0, 1.0, nil)
|
||||
ref.Addresses = []net.IP{net.ParseIP("1.2.3.4")}
|
||||
|
||||
tr := NewTraverser(&TraverserConfig{
|
||||
MaxDepth: 5,
|
||||
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||
})
|
||||
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||
return nil, nil
|
||||
})
|
||||
|
||||
ctx := context.Background()
|
||||
err := ref.Resolve(ctx, tr, nil, nil, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("Resolve with existing addresses: %v", err)
|
||||
}
|
||||
if ref.State != StateResolved {
|
||||
t.Errorf("expected StateResolved, got %v", ref.State)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReferralResolveCacheHit(t *testing.T) {
|
||||
ref := NewReferral("ns1.example.com.", dnsinternal.TypeA, ".", 0, 1.0, nil)
|
||||
|
||||
tr := NewTraverser(&TraverserConfig{
|
||||
MaxDepth: 5,
|
||||
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||
})
|
||||
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||
return nil, nil
|
||||
})
|
||||
|
||||
cache := NewInfoCache(nil)
|
||||
cache.StoreGlue("ns1.example.com.", []net.IP{net.ParseIP("1.2.3.4")})
|
||||
|
||||
ctx := context.Background()
|
||||
err := ref.Resolve(ctx, tr, cache, nil, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("Resolve cache hit: %v", err)
|
||||
}
|
||||
if ref.State != StateResolved {
|
||||
t.Errorf("expected StateResolved, got %v", ref.State)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReferralResolveCircular(t *testing.T) {
|
||||
ref := NewReferral("ns1.example.com.", dnsinternal.TypeA, ".", 0, 1.0, nil)
|
||||
|
||||
tr := NewTraverser(&TraverserConfig{
|
||||
MaxDepth: 5,
|
||||
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||
})
|
||||
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||
return nil, nil
|
||||
})
|
||||
|
||||
visited := map[string]bool{"ns1.example.com.": true}
|
||||
ctx := context.Background()
|
||||
err := ref.Resolve(ctx, tr, nil, visited, 0)
|
||||
if err == nil {
|
||||
t.Fatal("expected circular referral error")
|
||||
}
|
||||
if _, ok := err.(*CircularReferralError); !ok {
|
||||
t.Errorf("expected CircularReferralError, got %T: %v", err, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReferralResolveMaxDepth(t *testing.T) {
|
||||
ref := NewReferral("ns1.example.com.", dnsinternal.TypeA, ".", 0, 1.0, nil)
|
||||
|
||||
tr := NewTraverser(&TraverserConfig{
|
||||
MaxDepth: 5,
|
||||
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||
})
|
||||
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||
return nil, nil
|
||||
})
|
||||
|
||||
ctx := context.Background()
|
||||
err := ref.Resolve(ctx, tr, nil, nil, DefaultMaxDepth+1)
|
||||
if err == nil {
|
||||
t.Fatal("expected max depth error")
|
||||
}
|
||||
if _, ok := err.(*UnresolvableNameserverError); !ok {
|
||||
t.Errorf("expected UnresolvableNameserverError, got %T: %v", err, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReferralResolveWithAnswer(t *testing.T) {
|
||||
answerMsg := new(dns.Msg)
|
||||
answerMsg.SetReply(new(dns.Msg))
|
||||
answerMsg.Answer = append(answerMsg.Answer, &dns.A{
|
||||
Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300},
|
||||
A: net.ParseIP("1.2.3.4"),
|
||||
})
|
||||
|
||||
tr := NewTraverser(&TraverserConfig{
|
||||
MaxDepth: 5,
|
||||
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||
})
|
||||
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||
return answerMsg.Copy(), nil
|
||||
})
|
||||
|
||||
ref := NewReferral("ns1.example.com.", dnsinternal.TypeA, ".", 0, 1.0, nil)
|
||||
ctx := context.Background()
|
||||
err := ref.Resolve(ctx, tr, nil, nil, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("Resolve: %v", err)
|
||||
}
|
||||
if ref.State != StateResolved {
|
||||
t.Errorf("expected StateResolved, got %v", ref.State)
|
||||
}
|
||||
if len(ref.Addresses) == 0 {
|
||||
t.Error("expected addresses after resolution")
|
||||
}
|
||||
}
|
||||
|
||||
func TestReferralResolveNXDOMAIN(t *testing.T) {
|
||||
nxMsg := new(dns.Msg)
|
||||
nxMsg.Rcode = dns.RcodeNameError
|
||||
|
||||
tr := NewTraverser(&TraverserConfig{
|
||||
MaxDepth: 5,
|
||||
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||
})
|
||||
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||
return nxMsg.Copy(), nil
|
||||
})
|
||||
|
||||
ref := NewReferral("nonexistent.invalid.", dnsinternal.TypeA, ".", 0, 1.0, nil)
|
||||
ctx := context.Background()
|
||||
err := ref.Resolve(ctx, tr, nil, nil, 0)
|
||||
if err == nil {
|
||||
t.Fatal("expected error for NXDOMAIN")
|
||||
}
|
||||
if _, ok := err.(*UnresolvableNameserverError); !ok {
|
||||
t.Errorf("expected UnresolvableNameserverError, got %T: %v", err, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReferralResolveContextCancellation(t *testing.T) {
|
||||
tr := NewTraverser(&TraverserConfig{
|
||||
MaxDepth: 5,
|
||||
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||
})
|
||||
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||
refMsg := new(dns.Msg)
|
||||
refMsg.Rcode = dns.RcodeSuccess
|
||||
refMsg.Ns = append(refMsg.Ns, &dns.NS{
|
||||
Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS},
|
||||
Ns: "ns.example.com.",
|
||||
})
|
||||
refMsg.Extra = append(refMsg.Extra, &dns.A{
|
||||
Hdr: dns.RR_Header{Name: "ns.example.com.", Rrtype: dns.TypeA},
|
||||
A: net.ParseIP("1.2.3.4"),
|
||||
})
|
||||
return refMsg, nil
|
||||
})
|
||||
|
||||
ref := NewReferral("ns1.example.com.", dnsinternal.TypeA, ".", 0, 1.0, nil)
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel() // Cancel immediately
|
||||
|
||||
err := ref.Resolve(ctx, tr, nil, nil, 0)
|
||||
if err == nil {
|
||||
t.Fatal("expected error on cancelled context")
|
||||
}
|
||||
}
|
||||
|
||||
func TestReferralResolveReferralPath(t *testing.T) {
|
||||
// Test Resolve when it gets a referral response that pushes to stack
|
||||
referralMsg := new(dns.Msg)
|
||||
referralMsg.Rcode = dns.RcodeSuccess
|
||||
referralMsg.Ns = append(referralMsg.Ns, &dns.NS{
|
||||
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeNS, Class: dns.ClassINET},
|
||||
Ns: "ns1.example.com.",
|
||||
})
|
||||
referralMsg.Extra = append(referralMsg.Extra, &dns.A{
|
||||
Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET},
|
||||
A: net.ParseIP("1.2.3.4"),
|
||||
})
|
||||
|
||||
answerMsg := new(dns.Msg)
|
||||
answerMsg.SetReply(new(dns.Msg))
|
||||
answerMsg.Answer = append(answerMsg.Answer, &dns.A{
|
||||
Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300},
|
||||
A: net.ParseIP("5.6.7.8"),
|
||||
})
|
||||
|
||||
callCount := 0
|
||||
tr := NewTraverser(&TraverserConfig{
|
||||
MaxDepth: 5,
|
||||
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||
})
|
||||
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||
callCount++
|
||||
if callCount <= 1 {
|
||||
return referralMsg.Copy(), nil
|
||||
}
|
||||
return answerMsg.Copy(), nil
|
||||
})
|
||||
|
||||
ref := NewReferral("ns1.example.com.", dnsinternal.TypeA, ".", 0, 1.0, nil)
|
||||
ctx := context.Background()
|
||||
err := ref.Resolve(ctx, tr, nil, nil, 0)
|
||||
// May succeed or exhaust depending on referral loop
|
||||
t.Logf("Resolve referral path: err=%v, state=%v", err, ref.State)
|
||||
}
|
||||
|
||||
func TestResolutionStateStringUnknown(t *testing.T) {
|
||||
// Cover the default case of ResolutionState.String()
|
||||
unknown := ResolutionState(99)
|
||||
s := unknown.String()
|
||||
if s != "unknown" {
|
||||
t.Errorf("expected 'unknown' for invalid ResolutionState, got %q", s)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTraverserNonFastMode(t *testing.T) {
|
||||
answerResp := new(dns.Msg)
|
||||
answerResp.SetReply(new(dns.Msg))
|
||||
answerResp.Answer = append(answerResp.Answer, &dns.A{
|
||||
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300},
|
||||
A: net.ParseIP("1.2.3.4"),
|
||||
})
|
||||
|
||||
tr := NewTraverser(&TraverserConfig{
|
||||
MaxDepth: 5,
|
||||
QueryType: dnsinternal.TypeA,
|
||||
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||
Fast: false, // Non-fast mode
|
||||
})
|
||||
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||
return answerResp.Copy(), nil
|
||||
})
|
||||
|
||||
ctx := context.Background()
|
||||
results, err := tr.Traverse(ctx, "example.com")
|
||||
if err != nil {
|
||||
t.Fatalf("Traverse non-fast: %v", err)
|
||||
}
|
||||
if len(results) == 0 {
|
||||
t.Fatal("expected results")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIterativeQueryWithExchangeUsesConfig(t *testing.T) {
|
||||
answerMsg := new(dns.Msg)
|
||||
answerMsg.SetReply(new(dns.Msg))
|
||||
answerMsg.Answer = append(answerMsg.Answer, &dns.A{
|
||||
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300},
|
||||
A: net.ParseIP("1.2.3.4"),
|
||||
})
|
||||
|
||||
tr := NewTraverser(&TraverserConfig{
|
||||
MaxDepth: 5,
|
||||
QueryType: dnsinternal.TypeA,
|
||||
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||
QueryConfig: dnsinternal.DefaultQueryConfig(),
|
||||
})
|
||||
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||
return answerMsg.Copy(), nil
|
||||
})
|
||||
|
||||
ctx := context.Background()
|
||||
msg, err := tr.iterativeQueryWithExchange(ctx, net.ParseIP("198.41.0.4"), "example.com.", dnsinternal.TypeA)
|
||||
if err != nil {
|
||||
t.Fatalf("iterativeQueryWithExchange with config: %v", err)
|
||||
}
|
||||
if msg == nil {
|
||||
t.Fatal("expected non-nil response")
|
||||
}
|
||||
}
|
||||
|
||||
func TestTraverserReferralWithHooks(t *testing.T) {
|
||||
// Tests that hooks are called with IsResolve=true during ResolveNS sub-traversal
|
||||
answerMsg := new(dns.Msg)
|
||||
answerMsg.SetReply(new(dns.Msg))
|
||||
answerMsg.Answer = append(answerMsg.Answer, &dns.A{
|
||||
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300},
|
||||
A: net.ParseIP("1.2.3.4"),
|
||||
})
|
||||
|
||||
tr := NewTraverser(&TraverserConfig{
|
||||
MaxDepth: 5,
|
||||
QueryType: dnsinternal.TypeA,
|
||||
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||
})
|
||||
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||
return answerMsg.Copy(), nil
|
||||
})
|
||||
|
||||
var resolveEvents, progressEvents int
|
||||
tr.SetHooks(&TraverserHooks{
|
||||
OnEvent: func(e TraversalEvent) {
|
||||
if e.IsResolve {
|
||||
resolveEvents++
|
||||
} else {
|
||||
progressEvents++
|
||||
}
|
||||
},
|
||||
})
|
||||
|
||||
// Test directly via ResolveNS with hooks
|
||||
ctx := context.Background()
|
||||
addrs, err := tr.ResolveNS(ctx, "ns1.example.com.", nil, nil, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("ResolveNS: %v", err)
|
||||
}
|
||||
_ = addrs
|
||||
t.Logf("resolveEvents=%d progressEvents=%d", resolveEvents, progressEvents)
|
||||
}
|
||||
@@ -0,0 +1,259 @@
|
||||
package traverse
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
miekgdns "github.com/miekg/dns"
|
||||
)
|
||||
|
||||
// Status is the classification of one query outcome. The first eight values
|
||||
// come from decoded_query.rb / response.rb; noglue and loop are synthesised by
|
||||
// the resolve phase without sending a query.
|
||||
type Status string
|
||||
|
||||
const (
|
||||
StatusAnswered Status = "answered"
|
||||
StatusNoData Status = "nodata"
|
||||
StatusReferral Status = "referral"
|
||||
StatusRestart Status = "restart"
|
||||
StatusReferralLame Status = "referral_lame"
|
||||
StatusError Status = "error"
|
||||
StatusException Status = "exception"
|
||||
StatusCNAMELoop Status = "cname_loop"
|
||||
StatusNoGlue Status = "noglue"
|
||||
StatusLoop Status = "loop"
|
||||
)
|
||||
|
||||
// DecodedQuery classifies one DNS response (or network failure) against the
|
||||
// query that produced it, mirroring decoded_query.rb. Names are canonical
|
||||
// (lowercase, no trailing dot); Bailiwick "" means root.
|
||||
type DecodedQuery struct {
|
||||
Msg *miekgdns.Msg
|
||||
Err error
|
||||
|
||||
Qname string
|
||||
Qclass uint16
|
||||
Qtype uint16
|
||||
IP string
|
||||
Bailiwick string
|
||||
|
||||
Status Status
|
||||
Endname string
|
||||
// ChainTargets lists every CNAME target the in-message chain went
|
||||
// through (including the final unfollowed target when the chain leaves
|
||||
// the bailiwick); used for cross-response restart loop detection.
|
||||
ChainTargets []string
|
||||
|
||||
CacheableGood []miekgdns.RR
|
||||
CacheableBad []miekgdns.RR
|
||||
|
||||
AuthNS []miekgdns.RR
|
||||
AuthSOA []miekgdns.RR
|
||||
AuthOther []miekgdns.RR
|
||||
|
||||
Answers []miekgdns.RR
|
||||
AuthorityNames []string
|
||||
|
||||
ErrorMessage string
|
||||
ExceptionMessage string
|
||||
Warnings []string
|
||||
}
|
||||
|
||||
// NewDecodedQuery decodes and classifies a response. Pass err non-nil for a
|
||||
// network-level failure (dnstraverse's "exception"); msg is ignored then.
|
||||
func NewDecodedQuery(msg *miekgdns.Msg, err error, qname string, qclass, qtype uint16, ip, bailiwick string) *DecodedQuery {
|
||||
dq := &DecodedQuery{
|
||||
Msg: msg,
|
||||
Err: err,
|
||||
Qname: canonicalName(qname),
|
||||
Qclass: qclass,
|
||||
Qtype: qtype,
|
||||
IP: ip,
|
||||
Bailiwick: canonicalName(bailiwick),
|
||||
}
|
||||
dq.process()
|
||||
return dq
|
||||
}
|
||||
|
||||
func (dq *DecodedQuery) WarningsAdd(warnings ...string) {
|
||||
dq.Warnings = append(dq.Warnings, warnings...)
|
||||
}
|
||||
|
||||
// process implements the classification order of decoded_query.rb#process
|
||||
// exactly (the 7 steps in the design doc).
|
||||
func (dq *DecodedQuery) process() {
|
||||
if dq.Err == nil && dq.Msg == nil {
|
||||
dq.Err = fmt.Errorf("nil DNS response")
|
||||
}
|
||||
if dq.Err != nil {
|
||||
dq.Status = StatusException
|
||||
dq.ExceptionMessage = dq.Err.Error()
|
||||
return
|
||||
}
|
||||
dq.AuthNS, dq.AuthSOA, dq.AuthOther = msgAuthority(dq.Msg)
|
||||
dq.CacheableGood, dq.CacheableBad = msgCacheable(dq.Msg, dq.Bailiwick)
|
||||
endname, targets, ok := msgFollowCNAMEs(dq.Msg, dq.Qname, dq.Qtype, dq.Bailiwick)
|
||||
if !ok {
|
||||
dq.Status = StatusCNAMELoop
|
||||
return
|
||||
}
|
||||
dq.Endname = endname
|
||||
dq.ChainTargets = targets
|
||||
if dq.Msg.Rcode != miekgdns.RcodeSuccess {
|
||||
dq.Status = StatusError
|
||||
dq.ErrorMessage = rcodeErrorMessage(dq.Msg.Rcode)
|
||||
return
|
||||
}
|
||||
if answers := msgAnswers(dq.Msg, dq.Endname, dq.Qclass, dq.Qtype); len(answers) > 0 {
|
||||
dq.Answers = answers
|
||||
dq.Status = StatusAnswered
|
||||
return
|
||||
}
|
||||
if dq.Endname != dq.Qname {
|
||||
dq.Status = StatusRestart
|
||||
return
|
||||
}
|
||||
if len(dq.AuthSOA) > 0 || len(dq.AuthNS) == 0 {
|
||||
dq.Status = StatusNoData
|
||||
return
|
||||
}
|
||||
dq.Status = StatusReferral
|
||||
for _, rr := range dq.AuthNS {
|
||||
if ns, ok := rr.(*miekgdns.NS); ok {
|
||||
dq.AuthorityNames = append(dq.AuthorityNames, canonicalName(ns.Ns))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// rcodeErrorMessage renders the exact error strings of decoded_query.rb
|
||||
// process_error ("Format error" deliberately fixes the Ruby "Formate" typo —
|
||||
// documented deviation).
|
||||
func rcodeErrorMessage(rcode int) string {
|
||||
switch rcode {
|
||||
case miekgdns.RcodeFormatError:
|
||||
return "Format error (FORMERR)"
|
||||
case miekgdns.RcodeServerFailure:
|
||||
return "Server failure (SERVFAIL)"
|
||||
case miekgdns.RcodeNameError:
|
||||
return "No such domain (NXDOMAIN)"
|
||||
case miekgdns.RcodeNotImplemented:
|
||||
return "Not implemented (NOTIMP)"
|
||||
case miekgdns.RcodeRefused:
|
||||
return "Refused"
|
||||
default:
|
||||
if s, ok := miekgdns.RcodeToString[rcode]; ok {
|
||||
return s
|
||||
}
|
||||
return fmt.Sprintf("RCODE%d", rcode)
|
||||
}
|
||||
}
|
||||
|
||||
// insideBailiwick reports whether name is at or below the bailiwick zone:
|
||||
// bailiwick "" (root), equal fold, or name ends with "."+bailiwick.
|
||||
func insideBailiwick(name, bailiwick string) bool {
|
||||
bw := canonicalName(bailiwick)
|
||||
if bw == "" {
|
||||
return true
|
||||
}
|
||||
n := canonicalName(name)
|
||||
return n == bw || strings.HasSuffix(n, "."+bw)
|
||||
}
|
||||
|
||||
// msgAnswers returns the answer-section records matching qname/qclass/qtype
|
||||
// (message_utility.rb msg_answers?). qtype ANY matches every type.
|
||||
func msgAnswers(msg *miekgdns.Msg, qname string, qclass, qtype uint16) []miekgdns.RR {
|
||||
name := canonicalName(qname)
|
||||
any := qtype == miekgdns.TypeANY
|
||||
var out []miekgdns.RR
|
||||
for _, rr := range msg.Answer {
|
||||
h := rr.Header()
|
||||
if canonicalName(h.Name) == name && h.Class == qclass && (any || h.Rrtype == qtype) {
|
||||
out = append(out, rr)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// msgAuthority partitions the authority section into IN NS, IN SOA and other
|
||||
// records (message_utility.rb msg_authority).
|
||||
func msgAuthority(msg *miekgdns.Msg) (ns, soa, other []miekgdns.RR) {
|
||||
for _, rr := range msg.Ns {
|
||||
h := rr.Header()
|
||||
switch {
|
||||
case h.Rrtype == miekgdns.TypeNS && h.Class == miekgdns.ClassINET:
|
||||
ns = append(ns, rr)
|
||||
case h.Rrtype == miekgdns.TypeSOA && h.Class == miekgdns.ClassINET:
|
||||
soa = append(soa, rr)
|
||||
default:
|
||||
other = append(other, rr)
|
||||
}
|
||||
}
|
||||
return ns, soa, other
|
||||
}
|
||||
|
||||
// msgCacheable partitions ALL sections (answer, authority, additional — in
|
||||
// that order) into in-bailiwick records worth caching and out-of-bailiwick
|
||||
// records to discard. OPT pseudo-records are dropped entirely.
|
||||
func msgCacheable(msg *miekgdns.Msg, bailiwick string) (good, bad []miekgdns.RR) {
|
||||
for _, section := range [][]miekgdns.RR{msg.Answer, msg.Ns, msg.Extra} {
|
||||
for _, rr := range section {
|
||||
if rr.Header().Rrtype == miekgdns.TypeOPT {
|
||||
continue
|
||||
}
|
||||
if insideBailiwick(rr.Header().Name, bailiwick) {
|
||||
good = append(good, rr)
|
||||
} else {
|
||||
bad = append(bad, rr)
|
||||
}
|
||||
}
|
||||
}
|
||||
return good, bad
|
||||
}
|
||||
|
||||
// msgFollowCNAMEs follows a CNAME chain within one message and returns the
|
||||
// final name plus every target passed through (message_utility.rb
|
||||
// msg_follow_cnames). Following stops — the target is returned unfollowed —
|
||||
// as soon as the CURRENT owner name is not strictly below the bailiwick
|
||||
// (Ruby tests `name !~ /\.#{bailiwick}$/i`, so an owner exactly equal to the
|
||||
// bailiwick also stops the chain). An in-message loop returns ok=false
|
||||
// (cname_loop).
|
||||
func msgFollowCNAMEs(msg *miekgdns.Msg, qname string, qtype uint16, bailiwick string) (endname string, targets []string, ok bool) {
|
||||
name := canonicalName(qname)
|
||||
bw := canonicalName(bailiwick)
|
||||
seen := make(map[string]bool)
|
||||
for {
|
||||
seen[name] = true
|
||||
if len(msgAnswers(msg, name, miekgdns.ClassINET, qtype)) > 0 {
|
||||
return name, targets, true
|
||||
}
|
||||
cnames := msgAnswers(msg, name, miekgdns.ClassINET, miekgdns.TypeCNAME)
|
||||
if len(cnames) == 0 {
|
||||
return name, targets, true
|
||||
}
|
||||
cname, isCNAME := cnames[0].(*miekgdns.CNAME)
|
||||
if !isCNAME {
|
||||
return name, targets, true
|
||||
}
|
||||
target := canonicalName(cname.Target)
|
||||
targets = append(targets, target)
|
||||
if bw != "" && !strings.HasSuffix(name, "."+bw) {
|
||||
return target, targets, true
|
||||
}
|
||||
name = target
|
||||
if seen[name] {
|
||||
return "", targets, false
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// isLameReferral implements the response.rb lame rule: a referral is lame
|
||||
// unless the current bailiwick is root ("") or the new zone is STRICTLY
|
||||
// deeper than the current bailiwick (equal or sideways zones are lame).
|
||||
func isLameReferral(bailiwick, newBailiwick string) bool {
|
||||
bw := canonicalName(bailiwick)
|
||||
if bw == "" {
|
||||
return false
|
||||
}
|
||||
return !strings.HasSuffix(canonicalName(newBailiwick), "."+bw)
|
||||
}
|
||||
@@ -0,0 +1,359 @@
|
||||
package traverse
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"github.com/miekg/dns"
|
||||
)
|
||||
|
||||
func cnameRR(owner, target string) dns.RR {
|
||||
return &dns.CNAME{
|
||||
Hdr: dns.RR_Header{Name: dns.Fqdn(owner), Rrtype: dns.TypeCNAME, Class: dns.ClassINET},
|
||||
Target: dns.Fqdn(target),
|
||||
}
|
||||
}
|
||||
|
||||
func soaRR(zone string) dns.RR {
|
||||
return &dns.SOA{
|
||||
Hdr: dns.RR_Header{Name: dns.Fqdn(zone), Rrtype: dns.TypeSOA, Class: dns.ClassINET},
|
||||
Ns: dns.Fqdn("ns1." + zone),
|
||||
Mbox: dns.Fqdn("hostmaster." + zone),
|
||||
Serial: 1,
|
||||
Refresh: 3600, Retry: 600, Expire: 86400, Minttl: 300,
|
||||
}
|
||||
}
|
||||
|
||||
func newMsg(qname string, qtype uint16, rcode int) *dns.Msg {
|
||||
m := new(dns.Msg)
|
||||
m.SetQuestion(dns.Fqdn(qname), qtype)
|
||||
m.Response = true
|
||||
m.Rcode = rcode
|
||||
return m
|
||||
}
|
||||
|
||||
func decode(msg *dns.Msg, qname string, qtype uint16, bailiwick string) *DecodedQuery {
|
||||
return NewDecodedQuery(msg, nil, qname, dns.ClassINET, qtype, "192.0.2.1", bailiwick)
|
||||
}
|
||||
|
||||
func TestDecodeException(t *testing.T) {
|
||||
dq := NewDecodedQuery(nil, errors.New("network timeout"), "example.com", dns.ClassINET, dns.TypeA, "192.0.2.1", "com")
|
||||
if dq.Status != StatusException {
|
||||
t.Fatalf("status = %s, want exception", dq.Status)
|
||||
}
|
||||
if dq.ExceptionMessage != "network timeout" {
|
||||
t.Errorf("exception message = %q", dq.ExceptionMessage)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecodeNilMessageIsException(t *testing.T) {
|
||||
dq := NewDecodedQuery(nil, nil, "example.com", dns.ClassINET, dns.TypeA, "192.0.2.1", "com")
|
||||
if dq.Status != StatusException {
|
||||
t.Fatalf("status = %s, want exception", dq.Status)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecodeErrorMessages(t *testing.T) {
|
||||
tests := []struct {
|
||||
rcode int
|
||||
want string
|
||||
}{
|
||||
{dns.RcodeFormatError, "Format error (FORMERR)"},
|
||||
{dns.RcodeServerFailure, "Server failure (SERVFAIL)"},
|
||||
{dns.RcodeNameError, "No such domain (NXDOMAIN)"},
|
||||
{dns.RcodeNotImplemented, "Not implemented (NOTIMP)"},
|
||||
{dns.RcodeRefused, "Refused"},
|
||||
{dns.RcodeYXDomain, "YXDOMAIN"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
msg := newMsg("example.com", dns.TypeA, tt.rcode)
|
||||
dq := decode(msg, "example.com", dns.TypeA, "com")
|
||||
if dq.Status != StatusError {
|
||||
t.Errorf("rcode %d: status = %s, want error", tt.rcode, dq.Status)
|
||||
}
|
||||
if dq.ErrorMessage != tt.want {
|
||||
t.Errorf("rcode %d: message = %q, want %q", tt.rcode, dq.ErrorMessage, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecodeAnswered(t *testing.T) {
|
||||
msg := newMsg("example.com", dns.TypeA, dns.RcodeSuccess)
|
||||
msg.Answer = append(msg.Answer, aRR("example.com", "93.184.216.34"))
|
||||
dq := decode(msg, "example.com", dns.TypeA, "example.com")
|
||||
if dq.Status != StatusAnswered {
|
||||
t.Fatalf("status = %s, want answered", dq.Status)
|
||||
}
|
||||
if len(dq.Answers) != 1 {
|
||||
t.Errorf("answers = %v", dq.Answers)
|
||||
}
|
||||
if dq.Endname != "example.com" {
|
||||
t.Errorf("endname = %q", dq.Endname)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecodeAnsweredViaCNAMEChain(t *testing.T) {
|
||||
// In-bailiwick chain ends at a name that has the A answer.
|
||||
msg := newMsg("www.example.com", dns.TypeA, dns.RcodeSuccess)
|
||||
msg.Answer = append(msg.Answer,
|
||||
cnameRR("www.example.com", "web.example.com"),
|
||||
aRR("web.example.com", "93.184.216.34"),
|
||||
)
|
||||
dq := decode(msg, "www.example.com", dns.TypeA, "example.com")
|
||||
if dq.Status != StatusAnswered {
|
||||
t.Fatalf("status = %s, want answered", dq.Status)
|
||||
}
|
||||
if dq.Endname != "web.example.com" {
|
||||
t.Errorf("endname = %q, want web.example.com", dq.Endname)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecodeRestartOnOutOfBailiwickCNAME(t *testing.T) {
|
||||
msg := newMsg("www.example.com", dns.TypeA, dns.RcodeSuccess)
|
||||
msg.Answer = append(msg.Answer, cnameRR("www.example.com", "cdn.example.org"))
|
||||
dq := decode(msg, "www.example.com", dns.TypeA, "example.com")
|
||||
if dq.Status != StatusRestart {
|
||||
t.Fatalf("status = %s, want restart", dq.Status)
|
||||
}
|
||||
if dq.Endname != "cdn.example.org" {
|
||||
t.Errorf("endname = %q", dq.Endname)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecodeCNAMELoop(t *testing.T) {
|
||||
msg := newMsg("a.example.com", dns.TypeA, dns.RcodeSuccess)
|
||||
msg.Answer = append(msg.Answer,
|
||||
cnameRR("a.example.com", "b.example.com"),
|
||||
cnameRR("b.example.com", "a.example.com"),
|
||||
)
|
||||
dq := decode(msg, "a.example.com", dns.TypeA, "example.com")
|
||||
if dq.Status != StatusCNAMELoop {
|
||||
t.Fatalf("status = %s, want cname_loop", dq.Status)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecodeCNAMESelfLoop(t *testing.T) {
|
||||
msg := newMsg("a.example.com", dns.TypeA, dns.RcodeSuccess)
|
||||
msg.Answer = append(msg.Answer, cnameRR("a.example.com", "A.EXAMPLE.COM"))
|
||||
dq := decode(msg, "a.example.com", dns.TypeA, "example.com")
|
||||
if dq.Status != StatusCNAMELoop {
|
||||
t.Fatalf("status = %s, want cname_loop (case-insensitive)", dq.Status)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecodeCNAMEChainStopsAtOutOfBailiwickOwner(t *testing.T) {
|
||||
// Ruby stops following once the CURRENT owner leaves the bailiwick, so a
|
||||
// two-hop loop through an out-of-bailiwick owner is NOT cname_loop: the
|
||||
// unfollowed target equals the qname again, leaving endname == qname and
|
||||
// an empty authority — nodata.
|
||||
msg := newMsg("www.example.com", dns.TypeA, dns.RcodeSuccess)
|
||||
msg.Answer = append(msg.Answer,
|
||||
cnameRR("www.example.com", "a.example.org"),
|
||||
cnameRR("a.example.org", "www.example.com"),
|
||||
)
|
||||
dq := decode(msg, "www.example.com", dns.TypeA, "example.com")
|
||||
if dq.Status != StatusNoData {
|
||||
t.Fatalf("status = %s, want nodata", dq.Status)
|
||||
}
|
||||
if dq.Endname != "www.example.com" {
|
||||
t.Errorf("endname = %q, want www.example.com", dq.Endname)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecodeCNAMEOwnerEqualToBailiwickStopsChain(t *testing.T) {
|
||||
// Owner exactly equal to the bailiwick is NOT strictly inside it, so the
|
||||
// chain stops after one hop even though another CNAME exists.
|
||||
msg := newMsg("example.com", dns.TypeA, dns.RcodeSuccess)
|
||||
msg.Answer = append(msg.Answer,
|
||||
cnameRR("example.com", "a.example.com"),
|
||||
cnameRR("a.example.com", "b.example.com"),
|
||||
)
|
||||
dq := decode(msg, "example.com", dns.TypeA, "example.com")
|
||||
if dq.Status != StatusRestart {
|
||||
t.Fatalf("status = %s, want restart", dq.Status)
|
||||
}
|
||||
if dq.Endname != "a.example.com" {
|
||||
t.Errorf("endname = %q, want a.example.com (unfollowed target)", dq.Endname)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecodeQtypeCNAMEIsAnswered(t *testing.T) {
|
||||
// qtype=CNAME: the CNAME record IS the answer; the chain is never followed.
|
||||
msg := newMsg("www.example.com", dns.TypeCNAME, dns.RcodeSuccess)
|
||||
msg.Answer = append(msg.Answer,
|
||||
cnameRR("www.example.com", "web.example.com"),
|
||||
cnameRR("web.example.com", "www.example.com"),
|
||||
)
|
||||
dq := decode(msg, "www.example.com", dns.TypeCNAME, "example.com")
|
||||
if dq.Status != StatusAnswered {
|
||||
t.Fatalf("status = %s, want answered", dq.Status)
|
||||
}
|
||||
if dq.Endname != "www.example.com" {
|
||||
t.Errorf("endname = %q", dq.Endname)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecodeQtypeANYMatchesAnyAnswer(t *testing.T) {
|
||||
msg := newMsg("example.com", dns.TypeANY, dns.RcodeSuccess)
|
||||
msg.Answer = append(msg.Answer, cnameRR("example.com", "elsewhere.example.net"))
|
||||
dq := decode(msg, "example.com", dns.TypeANY, "example.com")
|
||||
if dq.Status != StatusAnswered {
|
||||
t.Fatalf("status = %s, want answered (ANY matches CNAME)", dq.Status)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecodeNoDataWithSOA(t *testing.T) {
|
||||
msg := newMsg("example.com", dns.TypeMX, dns.RcodeSuccess)
|
||||
msg.Ns = append(msg.Ns, soaRR("example.com"))
|
||||
dq := decode(msg, "example.com", dns.TypeMX, "example.com")
|
||||
if dq.Status != StatusNoData {
|
||||
t.Fatalf("status = %s, want nodata", dq.Status)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecodeNoDataEmptyAuthority(t *testing.T) {
|
||||
msg := newMsg("example.com", dns.TypeMX, dns.RcodeSuccess)
|
||||
dq := decode(msg, "example.com", dns.TypeMX, "example.com")
|
||||
if dq.Status != StatusNoData {
|
||||
t.Fatalf("status = %s, want nodata", dq.Status)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecodeNoDataSOAWinsOverNS(t *testing.T) {
|
||||
// SOA + NS in authority is a negative answer, not a referral.
|
||||
msg := newMsg("example.com", dns.TypeMX, dns.RcodeSuccess)
|
||||
msg.Ns = append(msg.Ns, soaRR("example.com"), nsRR("example.com", "ns1.example.com"))
|
||||
dq := decode(msg, "example.com", dns.TypeMX, "example.com")
|
||||
if dq.Status != StatusNoData {
|
||||
t.Fatalf("status = %s, want nodata", dq.Status)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecodeReferral(t *testing.T) {
|
||||
msg := newMsg("www.example.com", dns.TypeA, dns.RcodeSuccess)
|
||||
msg.Ns = append(msg.Ns,
|
||||
nsRR("example.com", "NS1.Example.COM"),
|
||||
nsRR("example.com", "ns2.example.net"),
|
||||
)
|
||||
msg.Extra = append(msg.Extra, aRR("ns1.example.com", "1.2.3.4"))
|
||||
dq := decode(msg, "www.example.com", dns.TypeA, "com")
|
||||
if dq.Status != StatusReferral {
|
||||
t.Fatalf("status = %s, want referral", dq.Status)
|
||||
}
|
||||
if len(dq.AuthorityNames) != 2 || dq.AuthorityNames[0] != "ns1.example.com" || dq.AuthorityNames[1] != "ns2.example.net" {
|
||||
t.Errorf("authority names = %v", dq.AuthorityNames)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecodeErrorBeatsAnswer(t *testing.T) {
|
||||
// rcode is checked before answers (step 3 before step 4).
|
||||
msg := newMsg("example.com", dns.TypeA, dns.RcodeServerFailure)
|
||||
msg.Answer = append(msg.Answer, aRR("example.com", "1.2.3.4"))
|
||||
dq := decode(msg, "example.com", dns.TypeA, "com")
|
||||
if dq.Status != StatusError {
|
||||
t.Fatalf("status = %s, want error", dq.Status)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecodeCNAMEFollowedIntoNXDOMAIN(t *testing.T) {
|
||||
// CNAME followed first (step 2), then rcode (step 3): NXDOMAIN after an
|
||||
// in-message CNAME is still an error, but the loop check ran first.
|
||||
msg := newMsg("www.example.com", dns.TypeA, dns.RcodeNameError)
|
||||
msg.Answer = append(msg.Answer, cnameRR("www.example.com", "gone.example.com"))
|
||||
dq := decode(msg, "www.example.com", dns.TypeA, "example.com")
|
||||
if dq.Status != StatusError {
|
||||
t.Fatalf("status = %s, want error", dq.Status)
|
||||
}
|
||||
if dq.Endname != "gone.example.com" {
|
||||
t.Errorf("endname = %q", dq.Endname)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecodeCacheablePartition(t *testing.T) {
|
||||
msg := newMsg("www.example.com", dns.TypeA, dns.RcodeSuccess)
|
||||
msg.Answer = append(msg.Answer, cnameRR("www.example.com", "cdn.example.org"))
|
||||
msg.Ns = append(msg.Ns, nsRR("example.org", "ns1.example.org"))
|
||||
msg.Extra = append(msg.Extra, aRR("ns1.example.org", "5.6.7.8"))
|
||||
opt := new(dns.OPT)
|
||||
opt.Hdr = dns.RR_Header{Name: ".", Rrtype: dns.TypeOPT}
|
||||
msg.Extra = append(msg.Extra, opt)
|
||||
|
||||
dq := decode(msg, "www.example.com", dns.TypeA, "example.com")
|
||||
if len(dq.CacheableGood) != 1 {
|
||||
t.Errorf("good = %v, want just the CNAME", dq.CacheableGood)
|
||||
}
|
||||
if len(dq.CacheableBad) != 2 {
|
||||
t.Errorf("bad = %v, want NS+A for example.org", dq.CacheableBad)
|
||||
}
|
||||
}
|
||||
|
||||
func TestInsideBailiwick(t *testing.T) {
|
||||
tests := []struct {
|
||||
name, bailiwick string
|
||||
want bool
|
||||
}{
|
||||
{"anything.example.com", "", true}, // root bailiwick
|
||||
{"anything.example.com", ".", true}, // root as dot
|
||||
{"example.com", "example.com", true}, // exact
|
||||
{"Example.COM", "example.com", true}, // exact, case fold
|
||||
{"www.example.com", "EXAMPLE.com", true}, // suffix, case fold
|
||||
{"a.b.example.com", "example.com", true}, // deep suffix
|
||||
{"badexample.com", "example.com", false}, // label boundary
|
||||
{"example.org", "example.com", false}, // sideways
|
||||
{"com", "example.com", false}, // shallower
|
||||
{"www.example.com.", "example.com", true}, // trailing dot
|
||||
}
|
||||
for _, tt := range tests {
|
||||
if got := insideBailiwick(tt.name, tt.bailiwick); got != tt.want {
|
||||
t.Errorf("insideBailiwick(%q, %q) = %v, want %v", tt.name, tt.bailiwick, got, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsLameReferral(t *testing.T) {
|
||||
tests := []struct {
|
||||
bailiwick, newBailiwick string
|
||||
want bool
|
||||
}{
|
||||
{"", "com", false}, // root bailiwick never lame
|
||||
{"", "", false}, // root to root
|
||||
{"com", "example.com", false}, // strictly deeper
|
||||
{"com", "a.b.example.com", false}, // much deeper
|
||||
{"COM", "example.com", false}, // case fold
|
||||
{"com", "com", true}, // equal zone is lame
|
||||
{"com", "", true}, // back to root is lame
|
||||
{"com", "org", true}, // sideways is lame
|
||||
{"example.com", "com", true}, // shallower is lame
|
||||
{"example.com", "badexample.com", true}, // label boundary
|
||||
{"example.com", "www.example.com", false}, // deeper
|
||||
}
|
||||
for _, tt := range tests {
|
||||
if got := isLameReferral(tt.bailiwick, tt.newBailiwick); got != tt.want {
|
||||
t.Errorf("isLameReferral(%q, %q) = %v, want %v", tt.bailiwick, tt.newBailiwick, got, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestMsgFollowCNAMEsNoChain(t *testing.T) {
|
||||
msg := newMsg("example.com", dns.TypeA, dns.RcodeSuccess)
|
||||
end, _, ok := msgFollowCNAMEs(msg, "Example.COM.", dns.TypeA, "com")
|
||||
if !ok || end != "example.com" {
|
||||
t.Errorf("end = %q ok=%v, want example.com true", end, ok)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMsgFollowCNAMEsRootBailiwickFollowsEverything(t *testing.T) {
|
||||
msg := newMsg("a.example.com", dns.TypeA, dns.RcodeSuccess)
|
||||
msg.Answer = append(msg.Answer,
|
||||
cnameRR("a.example.com", "b.example.org"),
|
||||
cnameRR("b.example.org", "c.example.net"),
|
||||
aRR("c.example.net", "1.2.3.4"),
|
||||
)
|
||||
end, targets, ok := msgFollowCNAMEs(msg, "a.example.com", dns.TypeA, "")
|
||||
if !ok || end != "c.example.net" {
|
||||
t.Errorf("end = %q ok=%v, want c.example.net true", end, ok)
|
||||
}
|
||||
if len(targets) != 2 || targets[0] != "b.example.org" || targets[1] != "c.example.net" {
|
||||
t.Errorf("chain targets = %v", targets)
|
||||
}
|
||||
}
|
||||
@@ -1,16 +1,67 @@
|
||||
package traverse
|
||||
|
||||
// EventStage mirrors the :stage symbols reported by traverser.rb's
|
||||
// report_progress: new/start/answer/resolve plus the fast-mode and
|
||||
// multi-childset variants.
|
||||
type EventStage int
|
||||
|
||||
const (
|
||||
EventStart EventStage = iota
|
||||
EventComplete
|
||||
// StageNew fires when a referral node is created (before processing).
|
||||
StageNew EventStage = iota
|
||||
// StageStart fires when a referral is popped for processing.
|
||||
StageStart
|
||||
// StageNewReferralSet fires once per extra childset when more than one
|
||||
// IP of a server produced children.
|
||||
StageNewReferralSet
|
||||
// StageNewFast fires instead of StageNew when fast mode already knows
|
||||
// this referral will be completed from the memo.
|
||||
StageNewFast
|
||||
// StageResolve fires after a resolve subtree's statistics are folded
|
||||
// into the referral (post-order, Ruby's :calc_resolve marker).
|
||||
StageResolve
|
||||
// StageAnswer fires after a referral's statistics are calculated
|
||||
// (post-order, Ruby's :calc_answer marker).
|
||||
StageAnswer
|
||||
// StageAnswerFast fires when fast mode replaced the referral with an
|
||||
// earlier completed one instead of processing it.
|
||||
StageAnswerFast
|
||||
)
|
||||
|
||||
func (s EventStage) String() string {
|
||||
switch s {
|
||||
case StageNew:
|
||||
return "new"
|
||||
case StageStart:
|
||||
return "start"
|
||||
case StageNewReferralSet:
|
||||
return "new_referral_set"
|
||||
case StageNewFast:
|
||||
return "new_fast"
|
||||
case StageResolve:
|
||||
return "resolve"
|
||||
case StageAnswer:
|
||||
return "answer"
|
||||
case StageAnswerFast:
|
||||
return "answer_fast"
|
||||
default:
|
||||
return "unknown"
|
||||
}
|
||||
}
|
||||
|
||||
// TraversalEvent is one progress callback. RefID/Status/IsResolve are
|
||||
// denormalised from Referral so renderers (CLI, web) need not walk the tree.
|
||||
type TraversalEvent struct {
|
||||
Stage EventStage
|
||||
Result TraversalResult
|
||||
Stage EventStage
|
||||
Referral *Referral
|
||||
RefID string
|
||||
// Status summarises the referral's outcome so far (see
|
||||
// Referral.OverallStatus); empty before any response arrives.
|
||||
Status Status
|
||||
// IsResolve is true for nodes inside a glue-resolution subtree.
|
||||
IsResolve bool
|
||||
// CompletedEarlier carries the refid of the earlier identical referral
|
||||
// on fast-mode events (StageNewFast / StageAnswerFast).
|
||||
CompletedEarlier string
|
||||
}
|
||||
|
||||
type EventHandler func(TraversalEvent)
|
||||
@@ -19,13 +70,16 @@ type TraverserHooks struct {
|
||||
OnEvent EventHandler
|
||||
}
|
||||
|
||||
func (h *TraverserHooks) emit(stage EventStage, result TraversalResult, isResolve bool) {
|
||||
if h == nil || h.OnEvent == nil {
|
||||
func (h *TraverserHooks) emit(stage EventStage, r *Referral, completedEarlier string) {
|
||||
if h == nil || h.OnEvent == nil || r == nil {
|
||||
return
|
||||
}
|
||||
h.OnEvent(TraversalEvent{
|
||||
Stage: stage,
|
||||
Result: result,
|
||||
IsResolve: isResolve,
|
||||
Stage: stage,
|
||||
Referral: r,
|
||||
RefID: r.RefID,
|
||||
Status: r.OverallStatus(),
|
||||
IsResolve: r.IsResolve(),
|
||||
CompletedEarlier: completedEarlier,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -1,52 +1,71 @@
|
||||
package traverse
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"testing"
|
||||
|
||||
"github.com/miekg/dns"
|
||||
)
|
||||
|
||||
func TestTraverserHooksEmitEvents(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: dns.TypeA, Class: dns.ClassINET, Ttl: 300},
|
||||
A: net.ParseIP("93.184.216.34"),
|
||||
})
|
||||
return m
|
||||
}()
|
||||
|
||||
var events []TraversalEvent
|
||||
hooks := &TraverserHooks{
|
||||
OnEvent: func(event TraversalEvent) {
|
||||
events = append(events, event)
|
||||
},
|
||||
func TestEventStageStrings(t *testing.T) {
|
||||
tests := map[EventStage]string{
|
||||
StageNew: "new",
|
||||
StageStart: "start",
|
||||
StageNewReferralSet: "new_referral_set",
|
||||
StageNewFast: "new_fast",
|
||||
StageResolve: "resolve",
|
||||
StageAnswer: "answer",
|
||||
StageAnswerFast: "answer_fast",
|
||||
EventStage(99): "unknown",
|
||||
}
|
||||
for stage, want := range tests {
|
||||
if got := stage.String(); got != want {
|
||||
t.Errorf("EventStage(%d).String() = %q, want %q", stage, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestHooksEmitNilSafe(t *testing.T) {
|
||||
var h *TraverserHooks
|
||||
h.emit(StageNew, newTestReferral("ns1.example.com", nil), "") // must not panic
|
||||
(&TraverserHooks{}).emit(StageNew, newTestReferral("ns1.example.com", nil), "")
|
||||
(&TraverserHooks{OnEvent: func(TraversalEvent) { t.Fatal("emitted for nil referral") }}).emit(StageNew, nil, "")
|
||||
}
|
||||
|
||||
func TestHooksEventSequenceSimpleAnswer(t *testing.T) {
|
||||
m := newMockExchange()
|
||||
m.on("198.41.0.4", "example.com", dns.TypeA, answerMsg(aRR("example.com", "9.9.9.9")))
|
||||
|
||||
tr := NewTraverser(&TraverserConfig{
|
||||
MaxDepth: 5,
|
||||
QueryType: dns.TypeA,
|
||||
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||
Hooks: hooks,
|
||||
})
|
||||
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||
return answerResp.Copy(), nil
|
||||
})
|
||||
var got []string
|
||||
cfg := testConfig(false)
|
||||
cfg.Hooks = &TraverserHooks{OnEvent: func(ev TraversalEvent) {
|
||||
got = append(got, ev.Stage.String()+":"+ev.RefID)
|
||||
}}
|
||||
runTraversal(t, cfg, m, "example.com")
|
||||
|
||||
_, err := tr.Traverse(context.Background(), "example.com")
|
||||
if err != nil {
|
||||
t.Fatalf("Traverse: %v", err)
|
||||
want := []string{"new:", "start:", "new:1", "start:1", "answer:1", "answer:"}
|
||||
if len(got) != len(want) {
|
||||
t.Fatalf("events = %v, want %v", got, want)
|
||||
}
|
||||
if len(events) < 2 {
|
||||
t.Fatalf("expected start and complete events, got %d", len(events))
|
||||
for i := range want {
|
||||
if got[i] != want[i] {
|
||||
t.Fatalf("events = %v, want %v", got, want)
|
||||
}
|
||||
}
|
||||
if events[0].Stage != EventStart {
|
||||
t.Fatalf("first event stage = %v, want start", events[0].Stage)
|
||||
}
|
||||
if events[1].Stage != EventComplete {
|
||||
t.Fatalf("second event stage = %v, want complete", events[1].Stage)
|
||||
}
|
||||
|
||||
func TestHooksEventCarriesStatus(t *testing.T) {
|
||||
m := newMockExchange()
|
||||
m.on("198.41.0.4", "example.com", dns.TypeA, answerMsg(aRR("example.com", "9.9.9.9")))
|
||||
|
||||
var answerStatus Status
|
||||
cfg := testConfig(false)
|
||||
cfg.Hooks = &TraverserHooks{OnEvent: func(ev TraversalEvent) {
|
||||
if ev.Stage == StageAnswer && ev.RefID == "1" {
|
||||
answerStatus = ev.Status
|
||||
}
|
||||
}}
|
||||
runTraversal(t, cfg, m, "example.com")
|
||||
if answerStatus != StatusAnswered {
|
||||
t.Errorf("answer event status = %q, want answered", answerStatus)
|
||||
}
|
||||
}
|
||||
|
||||
+560
-219
@@ -4,6 +4,7 @@ import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net"
|
||||
"sort"
|
||||
"strings"
|
||||
|
||||
"gitea.hansenits.com.au/hits/ExploreDNS/internal/dns"
|
||||
@@ -11,44 +12,87 @@ import (
|
||||
"golang.org/x/net/idna"
|
||||
)
|
||||
|
||||
type ResolutionState int
|
||||
// DefaultMaxDepth is the default maximum referral depth (non-zero refid
|
||||
// components) before a "Maxdepth N exceeded" exception is injected.
|
||||
const DefaultMaxDepth = 20
|
||||
|
||||
// ReferralStatus is the resolve-phase status of a Referral node itself
|
||||
// (referral.rb @status), distinct from the per-IP response statuses.
|
||||
type ReferralStatus string
|
||||
|
||||
const (
|
||||
StateUnresolved ResolutionState = iota
|
||||
StateResolving
|
||||
StateResolved
|
||||
RefStatusNormal ReferralStatus = "normal"
|
||||
RefStatusNoGlue ReferralStatus = "noglue"
|
||||
RefStatusLoop ReferralStatus = "loop"
|
||||
)
|
||||
|
||||
func (s ResolutionState) String() string {
|
||||
switch s {
|
||||
case StateUnresolved:
|
||||
return "unresolved"
|
||||
case StateResolving:
|
||||
return "resolving"
|
||||
case StateResolved:
|
||||
return "resolved"
|
||||
default:
|
||||
return "unknown"
|
||||
}
|
||||
// StatsEntry is one aggregated leaf statistic: the probability mass that
|
||||
// ended in Response's outcome at Referral (referral.rb @stats values).
|
||||
type StatsEntry struct {
|
||||
Key string
|
||||
Prob float64
|
||||
Response *ServerResponse
|
||||
Referral *Referral
|
||||
}
|
||||
|
||||
// Referral represents one referral to a specific server for qname/qclass/
|
||||
// qtype (referral.rb). The synthetic top node ("rootroot") has Server == ""
|
||||
// and is never displayed; its children are the root servers.
|
||||
type Referral struct {
|
||||
Name string
|
||||
Qtype uint16
|
||||
Qclass uint16
|
||||
Bailiwick string
|
||||
|
||||
Addresses []net.IP
|
||||
State ResolutionState
|
||||
|
||||
NSName string
|
||||
RefID string
|
||||
Parent *Referral
|
||||
Depth int
|
||||
Prob float64
|
||||
|
||||
Qname string
|
||||
Qclass uint16
|
||||
Qtype uint16
|
||||
// NSAType is the record type used to resolve nameserver addresses
|
||||
// (always A; the reference is IPv4-only for transport).
|
||||
NSAType uint16
|
||||
|
||||
// Server is the NS hostname this referral queries ("" for rootroot).
|
||||
Server string
|
||||
// ServerIPs is nil when the server still needs resolving. After a
|
||||
// resolve it may also contain "key:..." pseudo entries carrying the
|
||||
// probability of failed resolutions.
|
||||
ServerIPs []string
|
||||
Bailiwick string
|
||||
ParentIP string
|
||||
|
||||
InfoCache *InfoCache
|
||||
Status ReferralStatus
|
||||
|
||||
// Responses holds the classified response for each real IP queried.
|
||||
Responses map[string]*ServerResponse
|
||||
// Children holds child referrals keyed by the parent IP that produced
|
||||
// them ("rootroot" for the synthetic top node).
|
||||
Children map[string][]*Referral
|
||||
// Resolves is the glue-resolution subtree (refid ".0." components).
|
||||
Resolves []*Referral
|
||||
|
||||
ServerWeights map[string]float64
|
||||
Warnings []string
|
||||
|
||||
// Stats is the post-order aggregation of leaf outcomes below this node.
|
||||
Stats map[string]*StatsEntry
|
||||
// StatsResolve aggregates the outcomes of the resolve subtree.
|
||||
StatsResolve map[string]*StatsEntry
|
||||
|
||||
// ReplacedBy points at the earlier completed referral that fast mode
|
||||
// substituted for this node.
|
||||
ReplacedBy *Referral
|
||||
|
||||
// summaryStats memoises SummaryStats() (referral.rb summary_stats).
|
||||
summaryStats *SummaryStats
|
||||
|
||||
client *dns.Client
|
||||
maxdepth int
|
||||
referralResolution bool
|
||||
processed bool
|
||||
calculated bool
|
||||
}
|
||||
|
||||
// idnaLookup is the IDN lookup profile used to convert internationalised domain
|
||||
// names (unicode labels) to their ACE/punycode equivalents before querying.
|
||||
// idnaLookup is the IDN lookup profile used to convert internationalised
|
||||
// domain names (unicode labels) to their ACE/punycode equivalents.
|
||||
var idnaLookup = idna.New(
|
||||
idna.MapForLookup(),
|
||||
idna.BidiRule(),
|
||||
@@ -56,9 +100,8 @@ var idnaLookup = idna.New(
|
||||
)
|
||||
|
||||
// toASCII converts a domain name that may contain unicode labels to its
|
||||
// punycode (ACE) representation. Pure-ASCII names are returned unchanged.
|
||||
// On conversion errors the original name is returned so the caller can still
|
||||
// attempt a query (the server will reject it if truly invalid).
|
||||
// punycode (ACE) representation. On conversion errors the original name is
|
||||
// returned so the caller can still attempt a query.
|
||||
func toASCII(name string) string {
|
||||
if name == "" || name == "." {
|
||||
return name
|
||||
@@ -70,210 +113,508 @@ func toASCII(name string) string {
|
||||
return ascii
|
||||
}
|
||||
|
||||
func NewReferral(name string, qtype uint16, bailiwick string, depth int, prob float64, parent *Referral) *Referral {
|
||||
return &Referral{
|
||||
Name: miekgdns.Fqdn(strings.ToLower(toASCII(name))),
|
||||
Qtype: qtype,
|
||||
Qclass: miekgdns.ClassINET,
|
||||
Bailiwick: miekgdns.Fqdn(strings.ToLower(toASCII(bailiwick))),
|
||||
Depth: depth,
|
||||
Prob: prob,
|
||||
Parent: parent,
|
||||
State: StateUnresolved,
|
||||
// referralArgs are the per-child overrides for makeReferral; zero values
|
||||
// inherit from the parent (referral.rb make_referral merge semantics).
|
||||
type referralArgs struct {
|
||||
qname string
|
||||
qtype uint16
|
||||
server string
|
||||
serverIPs []string
|
||||
bailiwick string
|
||||
infoCache *InfoCache
|
||||
refid string
|
||||
parentIP string
|
||||
referralResolution bool
|
||||
}
|
||||
|
||||
func (r *Referral) makeReferral(a referralArgs) *Referral {
|
||||
child := &Referral{
|
||||
RefID: a.refid,
|
||||
Parent: r,
|
||||
Qname: r.Qname,
|
||||
Qclass: r.Qclass,
|
||||
Qtype: r.Qtype,
|
||||
NSAType: r.NSAType,
|
||||
Server: canonicalName(a.server),
|
||||
ServerIPs: a.serverIPs,
|
||||
Bailiwick: canonicalName(a.bailiwick),
|
||||
ParentIP: a.parentIP,
|
||||
InfoCache: r.InfoCache,
|
||||
Status: RefStatusNormal,
|
||||
Responses: make(map[string]*ServerResponse),
|
||||
Children: make(map[string][]*Referral),
|
||||
ServerWeights: make(map[string]float64),
|
||||
client: r.client,
|
||||
maxdepth: r.maxdepth,
|
||||
referralResolution: a.referralResolution || r.referralResolution,
|
||||
}
|
||||
}
|
||||
|
||||
func (r *Referral) InBailiwick(name string) bool {
|
||||
if r.Bailiwick == "" || r.Bailiwick == "." {
|
||||
return true
|
||||
if a.qname != "" {
|
||||
child.Qname = canonicalName(a.qname)
|
||||
}
|
||||
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
|
||||
if a.qtype != 0 {
|
||||
child.Qtype = a.qtype
|
||||
}
|
||||
}
|
||||
|
||||
type CircularReferralError struct {
|
||||
Name string
|
||||
Chain []string
|
||||
}
|
||||
|
||||
func (e *CircularReferralError) Error() string {
|
||||
return fmt.Sprintf("circular referral detected for %s: %v", e.Name, e.Chain)
|
||||
}
|
||||
|
||||
type UnresolvableNameserverError struct {
|
||||
Name string
|
||||
Reason string
|
||||
}
|
||||
|
||||
func (e *UnresolvableNameserverError) Error() string {
|
||||
return fmt.Sprintf("unresolvable nameserver %s: %s", e.Name, e.Reason)
|
||||
}
|
||||
|
||||
func (r *Referral) Resolve(ctx context.Context, traverser *Traverser, cache *InfoCache, visited map[string]bool, depth int) error {
|
||||
if r.HasAddresses() {
|
||||
r.State = StateResolved
|
||||
return nil
|
||||
if a.infoCache != nil {
|
||||
child.InfoCache = a.infoCache
|
||||
}
|
||||
|
||||
if cache != nil {
|
||||
if addrs := cache.LookupGlue(r.Name); len(addrs) > 0 {
|
||||
r.Addresses = addrs
|
||||
r.State = StateResolved
|
||||
return nil
|
||||
// serverweight = 1/len(ips) per IP when the addresses are known.
|
||||
if child.ServerIPs != nil {
|
||||
for _, ip := range child.ServerIPs {
|
||||
child.ServerWeights[ip] = 1.0 / float64(len(child.ServerIPs))
|
||||
}
|
||||
}
|
||||
return child
|
||||
}
|
||||
|
||||
if visited != nil {
|
||||
if visited[r.Name] {
|
||||
return &CircularReferralError{
|
||||
Name: r.Name,
|
||||
Chain: getVisitedNames(visited),
|
||||
// IsRootRoot reports whether this is the synthetic top node representing an
|
||||
// automatic referral to the root servers.
|
||||
func (r *Referral) IsRootRoot() bool {
|
||||
return r.Server == ""
|
||||
}
|
||||
|
||||
// IsResolve reports whether this node is part of a glue-resolution subtree.
|
||||
func (r *Referral) IsResolve() bool {
|
||||
return r.referralResolution
|
||||
}
|
||||
|
||||
// Resolved reports whether the server addresses are known (rootroot is
|
||||
// always resolved).
|
||||
func (r *Referral) Resolved() bool {
|
||||
return r.IsRootRoot() || r.ServerIPs != nil
|
||||
}
|
||||
|
||||
// Depth counts the non-zero refid components; resolve subtrees (".0.") do
|
||||
// not count against the depth limit.
|
||||
func (r *Referral) Depth() int {
|
||||
return refidDepth(r.RefID)
|
||||
}
|
||||
|
||||
func refidDepth(refid string) int {
|
||||
if refid == "" {
|
||||
return 0
|
||||
}
|
||||
n := 0
|
||||
for _, part := range strings.Split(refid, ".") {
|
||||
if part != "0" {
|
||||
n++
|
||||
}
|
||||
}
|
||||
return n
|
||||
}
|
||||
|
||||
// IPsAsArray returns the real IP addresses known for this referral,
|
||||
// excluding "key:" pseudo entries.
|
||||
func (r *Referral) IPsAsArray() []string {
|
||||
var out []string
|
||||
for _, ip := range r.ServerIPs {
|
||||
if !strings.HasPrefix(ip, "key:") {
|
||||
out = append(out, ip)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// TxtIPsVerbose renders the per-IP weights, sorted, e.g.
|
||||
// "50.0%=1.2.3.4,50.0%=noglue:1.2.3.4" (referral.rb txt_ips_verbose). It is
|
||||
// part of the fast-mode memo key.
|
||||
func (r *Referral) TxtIPsVerbose() string {
|
||||
if r.ServerIPs == nil {
|
||||
return ""
|
||||
}
|
||||
parts := make([]string, 0, len(r.ServerIPs))
|
||||
for _, ip := range r.ServerIPs {
|
||||
label := ip
|
||||
if rest, ok := strings.CutPrefix(ip, "key:"); ok {
|
||||
// keep the first two colon-separated fields, like Ruby's
|
||||
// /^key:([^:]+(:[^:]*)?)/ capture.
|
||||
fields := strings.SplitN(rest, ":", 3)
|
||||
if len(fields) > 2 {
|
||||
fields = fields[:2]
|
||||
}
|
||||
label = strings.Join(fields, ":")
|
||||
}
|
||||
parts = append(parts, fmt.Sprintf("%.1f%%=%s", 100*r.ServerWeights[ip], label))
|
||||
}
|
||||
sort.Strings(parts)
|
||||
return strings.Join(parts, ",")
|
||||
}
|
||||
|
||||
// TxtIPs renders the addresses for progress display; failed-resolve pseudo
|
||||
// entries render as their response description (referral.rb txt_ips).
|
||||
func (r *Referral) TxtIPs() string {
|
||||
if r.ServerIPs == nil {
|
||||
return ""
|
||||
}
|
||||
parts := make([]string, 0, len(r.ServerIPs))
|
||||
for _, ip := range r.ServerIPs {
|
||||
if strings.HasPrefix(ip, "key:") {
|
||||
if e, ok := r.StatsResolve[ip]; ok && e.Response != nil {
|
||||
parts = append(parts, e.Response.String())
|
||||
continue
|
||||
}
|
||||
}
|
||||
visited[r.Name] = true
|
||||
}
|
||||
|
||||
if depth > DefaultMaxDepth {
|
||||
return &UnresolvableNameserverError{
|
||||
Name: r.Name,
|
||||
Reason: "max depth exceeded",
|
||||
}
|
||||
}
|
||||
|
||||
roots, err := traverser.discoverRoots(ctx)
|
||||
if err != nil {
|
||||
return fmt.Errorf("root discovery: %w", err)
|
||||
}
|
||||
|
||||
initial := NewReferral(r.Name, dns.TypeA, ".", 0, 1.0, nil)
|
||||
initial.Addresses = roots
|
||||
initial.State = StateResolved
|
||||
|
||||
stack := NewStack(DefaultMaxDepth)
|
||||
stack.Push(initial)
|
||||
|
||||
traversalCache := NewInfoCache(nil)
|
||||
if visited != nil {
|
||||
for name := range visited {
|
||||
traversalCache.StoreGlue(name, []net.IP{})
|
||||
}
|
||||
}
|
||||
|
||||
var lastErr error
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return fmt.Errorf("resolution cancelled: %w", ctx.Err())
|
||||
default:
|
||||
}
|
||||
|
||||
ref := stack.Pop()
|
||||
if ref == nil {
|
||||
break
|
||||
}
|
||||
|
||||
cacheForStep := traversalCache
|
||||
if ref.Parent != nil {
|
||||
cacheForStep = traversalCache.Child()
|
||||
}
|
||||
|
||||
resp := traverser.processReferral(ctx, ref, cacheForStep)
|
||||
|
||||
if resp.Type == RespAnswer {
|
||||
if len(resp.Decoded.Answers) > 0 {
|
||||
var addrs []net.IP
|
||||
for _, rr := range resp.Decoded.Answers {
|
||||
if a, ok := rr.(*miekgdns.A); ok {
|
||||
addrs = append(addrs, a.A)
|
||||
}
|
||||
if aaaa, ok := rr.(*miekgdns.AAAA); ok {
|
||||
addrs = append(addrs, aaaa.AAAA)
|
||||
}
|
||||
}
|
||||
if len(addrs) > 0 {
|
||||
r.Addresses = addrs
|
||||
r.State = StateResolved
|
||||
if cache != nil {
|
||||
cache.StoreGlue(r.Name, addrs)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if resp.Type == RespNXDOMAIN {
|
||||
lastErr = &UnresolvableNameserverError{
|
||||
Name: r.Name,
|
||||
Reason: "NXDOMAIN",
|
||||
}
|
||||
break
|
||||
}
|
||||
|
||||
if resp.Type == RespSERVFAIL || resp.Type == RespError {
|
||||
lastErr = fmt.Errorf("server error resolving %s: %s", r.Name, resp.Type)
|
||||
continue
|
||||
}
|
||||
|
||||
if resp.Type == RespReferral {
|
||||
children := resp.ChildReferrals()
|
||||
for _, child := range children {
|
||||
// Only skip visited names when they have no addresses; if glue
|
||||
// was included in the referral response we still need to query
|
||||
// that child to get the authoritative answer.
|
||||
if visited != nil && visited[child.Name] && !child.HasAddresses() {
|
||||
continue
|
||||
}
|
||||
if !stack.Push(child) {
|
||||
lastErr = &UnresolvableNameserverError{
|
||||
Name: r.Name,
|
||||
Reason: "max depth exceeded during resolution",
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if lastErr != nil {
|
||||
return lastErr
|
||||
}
|
||||
|
||||
return &UnresolvableNameserverError{
|
||||
Name: r.Name,
|
||||
Reason: "resolution exhausted without answer",
|
||||
parts = append(parts, ip)
|
||||
}
|
||||
sort.Strings(parts)
|
||||
return strings.Join(parts, ",")
|
||||
}
|
||||
|
||||
func getVisitedNames(visited map[string]bool) []string {
|
||||
var names []string
|
||||
for name := range visited {
|
||||
names = append(names, name)
|
||||
}
|
||||
return names
|
||||
func (r *Referral) String() string {
|
||||
return fmt.Sprintf("%s [%s/%s/%s] server=%s server_ips=%s bailiwick=%s",
|
||||
r.RefID, r.Qname, ClassToString(r.Qclass), TypeToString(r.Qtype),
|
||||
r.Server, r.TxtIPs(), r.Bailiwick)
|
||||
}
|
||||
|
||||
// IsNameInChain reports whether name appears anywhere in this referral's ancestor
|
||||
// chain, including this referral itself. Used for CNAME loop detection.
|
||||
func (r *Referral) IsNameInChain(name string) bool {
|
||||
n := miekgdns.Fqdn(strings.ToLower(name))
|
||||
curr := r
|
||||
for curr != nil {
|
||||
if curr.Name == n {
|
||||
// OverallStatus summarises the node's outcome for event consumers: a resolve
|
||||
// dead end (noglue/loop), the shared status of every per-IP response, or
|
||||
// "mixed" when the responses disagree ("" before anything was queried).
|
||||
func (r *Referral) OverallStatus() Status {
|
||||
switch r.Status {
|
||||
case RefStatusNoGlue:
|
||||
return StatusNoGlue
|
||||
case RefStatusLoop:
|
||||
return StatusLoop
|
||||
}
|
||||
var s Status
|
||||
for _, resp := range r.Responses {
|
||||
if s == "" {
|
||||
s = resp.Status
|
||||
} else if s != resp.Status {
|
||||
return "mixed"
|
||||
}
|
||||
}
|
||||
return s
|
||||
}
|
||||
|
||||
// StatsList returns the aggregated leaf statistics sorted by stats key.
|
||||
func (r *Referral) StatsList() []*StatsEntry {
|
||||
out := make([]*StatsEntry, 0, len(r.Stats))
|
||||
for _, e := range r.Stats {
|
||||
out = append(out, e)
|
||||
}
|
||||
sort.Slice(out, func(i, j int) bool { return out[i].Key < out[j].Key })
|
||||
return out
|
||||
}
|
||||
|
||||
// isNoGlue reports a dead end: the server is inside the current bailiwick,
|
||||
// so its address should have come as glue, but none was provided and no
|
||||
// deeper zone can name it (referral.rb noglue?).
|
||||
func (r *Referral) isNoGlue() bool {
|
||||
return r.ServerIPs == nil && insideBailiwick(r.Server, r.Bailiwick)
|
||||
}
|
||||
|
||||
// isLoop reports a resolve loop: an ancestor referral asks the same
|
||||
// qname/qclass/qtype of the same still-unresolved server (referral.rb loop?),
|
||||
// e.g. b NS c.d while d NS a.b.
|
||||
func (r *Referral) isLoop() bool {
|
||||
if r.ServerIPs != nil {
|
||||
return false
|
||||
}
|
||||
for p := r.Parent; p != nil; p = p.Parent {
|
||||
if p.Qname == r.Qname && p.Qclass == r.Qclass && p.Qtype == r.Qtype &&
|
||||
p.Server == r.Server && p.ServerIPs == nil {
|
||||
return true
|
||||
}
|
||||
curr = curr.Parent
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// chainHasQuery reports whether this referral or any ancestor already asks
|
||||
// qname with the same qclass/qtype. Used to stop cross-response CNAME chains
|
||||
// (restart loops) that Ruby only catches via the depth limit.
|
||||
func (r *Referral) chainHasQuery(qname string) bool {
|
||||
name := canonicalName(qname)
|
||||
for p := r; p != nil; p = p.Parent {
|
||||
if p.Qname == name && p.Qclass == r.Qclass && p.Qtype == r.Qtype {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// resolve turns an address-less referral into either a dead end (noglue/
|
||||
// loop) or a resolve subtree querying A <server> from this branch's cache
|
||||
// (referral.rb resolve). It returns the referrals to process.
|
||||
func (r *Referral) resolve() ([]*Referral, error) {
|
||||
if r.isNoGlue() {
|
||||
r.Status = RefStatusNoGlue
|
||||
return nil, nil
|
||||
}
|
||||
if r.isLoop() {
|
||||
r.Status = RefStatusLoop
|
||||
return nil, nil
|
||||
}
|
||||
starters, newbailiwick, err := r.InfoCache.GetStartServers(r.Server)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for i, st := range starters {
|
||||
child := r.makeReferral(referralArgs{
|
||||
qname: r.Server,
|
||||
qtype: r.NSAType,
|
||||
server: st.Name,
|
||||
serverIPs: st.IPs,
|
||||
bailiwick: newbailiwick,
|
||||
refid: fmt.Sprintf("%s.0.%d", r.RefID, i+1),
|
||||
referralResolution: true,
|
||||
})
|
||||
r.Resolves = append(r.Resolves, child)
|
||||
}
|
||||
return r.Resolves, nil
|
||||
}
|
||||
|
||||
// resolveCalculate folds the resolve subtree's statistics into per-IP server
|
||||
// weights (referral.rb resolve_calculate): each answered leaf distributes
|
||||
// its probability evenly across the A records it returned; every other leaf
|
||||
// keeps its probability under its "key:" stats key so failures surface in
|
||||
// the results.
|
||||
func (r *Referral) resolveCalculate() {
|
||||
r.StatsResolve = make(map[string]*StatsEntry)
|
||||
switch r.Status {
|
||||
case RefStatusNoGlue:
|
||||
resp := NewNoGlueResponse(r.Qname, r.Qclass, r.Qtype, r.ParentIP, r.Server, r.Bailiwick)
|
||||
key := resp.StatsKey()
|
||||
r.StatsResolve[key] = &StatsEntry{Key: key, Prob: 1.0, Response: resp, Referral: r}
|
||||
case RefStatusLoop:
|
||||
resp := NewLoopResponse(r.Qname, r.Qclass, r.Qtype, r.ParentIP, r.Server, r.Bailiwick)
|
||||
key := resp.StatsKey()
|
||||
r.StatsResolve[key] = &StatsEntry{Key: key, Prob: 1.0, Response: resp, Referral: r}
|
||||
default:
|
||||
statsCalculateChildren(r.StatsResolve, r.Resolves, 1.0)
|
||||
}
|
||||
|
||||
r.ServerWeights = make(map[string]float64)
|
||||
r.ServerIPs = []string{}
|
||||
keys := make([]string, 0, len(r.StatsResolve))
|
||||
for key := range r.StatsResolve {
|
||||
keys = append(keys, key)
|
||||
}
|
||||
sort.Strings(keys)
|
||||
addWeight := func(ip string, prob float64) {
|
||||
if _, ok := r.ServerWeights[ip]; !ok {
|
||||
r.ServerIPs = append(r.ServerIPs, ip)
|
||||
}
|
||||
r.ServerWeights[ip] += prob
|
||||
}
|
||||
for _, key := range keys {
|
||||
data := r.StatsResolve[key]
|
||||
if data.Response.Status == StatusAnswered {
|
||||
var addrs []string
|
||||
for _, rr := range data.Response.DQ.Answers {
|
||||
if a, ok := rr.(*miekgdns.A); ok {
|
||||
addrs = append(addrs, a.A.String())
|
||||
}
|
||||
}
|
||||
for _, addr := range addrs {
|
||||
addWeight(addr, data.Prob/float64(len(addrs)))
|
||||
}
|
||||
if len(addrs) == 0 {
|
||||
// answered but no A records (e.g. AAAA-only): carry the
|
||||
// probability as a failure key so mass is not lost.
|
||||
addWeight(key, data.Prob)
|
||||
}
|
||||
} else {
|
||||
addWeight(key, data.Prob)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// statsCalculateChildren merges the children's statistics into stats with an
|
||||
// equal split of weight among them (referral.rb stats_calculate_children).
|
||||
func statsCalculateChildren(stats map[string]*StatsEntry, children []*Referral, weight float64) {
|
||||
if len(children) == 0 {
|
||||
return
|
||||
}
|
||||
percent := (1.0 / float64(len(children))) * weight
|
||||
for _, child := range children {
|
||||
for key, data := range child.Stats {
|
||||
if e, ok := stats[key]; ok {
|
||||
e.Prob += data.Prob * percent
|
||||
} else {
|
||||
stats[key] = &StatsEntry{
|
||||
Key: key,
|
||||
Prob: data.Prob * percent,
|
||||
Response: data.Response,
|
||||
Referral: data.Referral,
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// answerCalculate computes this node's aggregated statistics from its
|
||||
// children, responses and resolve failures (referral.rb answer_calculate).
|
||||
// Unlike the Ruby source, duplicate stats keys across a referral's IPs merge
|
||||
// by summing probability (Ruby computes the sum then discards it — a source
|
||||
// bug that breaks the probabilities-sum-to-1 invariant).
|
||||
func (r *Referral) answerCalculate() {
|
||||
r.Stats = make(map[string]*StatsEntry)
|
||||
if r.IsRootRoot() {
|
||||
statsCalculateChildren(r.Stats, r.Children["rootroot"], 1.0)
|
||||
r.calculated = true
|
||||
return
|
||||
}
|
||||
for _, ip := range r.ServerIPs {
|
||||
serverweight := r.ServerWeights[ip]
|
||||
if strings.HasPrefix(ip, "key:") {
|
||||
// resolve failed for some reason - copy the resolve statistics
|
||||
src := r.StatsResolve[ip]
|
||||
if e, ok := r.Stats[ip]; ok {
|
||||
e.Prob += src.Prob
|
||||
} else {
|
||||
r.Stats[ip] = &StatsEntry{Key: ip, Prob: src.Prob, Response: src.Response, Referral: src.Referral}
|
||||
}
|
||||
continue
|
||||
}
|
||||
if children := r.Children[ip]; len(children) > 0 {
|
||||
statsCalculateChildren(r.Stats, children, serverweight)
|
||||
continue
|
||||
}
|
||||
resp := r.Responses[ip]
|
||||
if resp == nil {
|
||||
continue
|
||||
}
|
||||
key := resp.StatsKey()
|
||||
if e, ok := r.Stats[key]; ok {
|
||||
e.Prob += serverweight
|
||||
} else {
|
||||
r.Stats[key] = &StatsEntry{Key: key, Prob: serverweight, Response: resp, Referral: r}
|
||||
}
|
||||
}
|
||||
r.calculated = true
|
||||
}
|
||||
|
||||
// process queries every real IP of this referral through the packet cache,
|
||||
// classifies each response, and creates one child per NS name (including
|
||||
// glueless ones) for referral/restart statuses (referral.rb process/
|
||||
// process_normal). It returns one set of children per IP that produced any.
|
||||
func (r *Referral) process(ctx context.Context) ([][]*Referral, error) {
|
||||
r.processed = true
|
||||
if r.IsRootRoot() {
|
||||
children, err := r.processAddRoots()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return [][]*Referral{children}, nil
|
||||
}
|
||||
|
||||
// Phase one: query and classify, counting childsets so refids can grow
|
||||
// an extra childset digit when more than one IP produces children.
|
||||
childsets := 0
|
||||
var order []string
|
||||
for _, ip := range r.ServerIPs {
|
||||
if strings.HasPrefix(ip, "key:") {
|
||||
continue
|
||||
}
|
||||
var dq *DecodedQuery
|
||||
if r.Depth() >= r.maxdepth {
|
||||
err := fmt.Errorf("Maxdepth %d exceeded", r.maxdepth)
|
||||
dq = NewDecodedQuery(nil, err, r.Qname, r.Qclass, r.Qtype, ip, r.Bailiwick)
|
||||
} else {
|
||||
msg, warnings, err := r.client.Query(ctx, net.ParseIP(ip), r.Qname, r.Qtype)
|
||||
dq = NewDecodedQuery(msg, err, r.Qname, r.Qclass, r.Qtype, ip, r.Bailiwick)
|
||||
dq.WarningsAdd(warnings...)
|
||||
}
|
||||
resp, err := NewServerResponse(dq, r.Server, r.ParentIP, r.InfoCache)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if resp.Status == StatusRestart {
|
||||
// Cross-response CNAME loop: any target in the chain that we
|
||||
// (or an ancestor) are already querying is a dead end.
|
||||
for _, target := range dq.ChainTargets {
|
||||
if r.chainHasQuery(target) {
|
||||
resp.Status = StatusCNAMELoop
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
r.Warnings = append(r.Warnings, dq.Warnings...)
|
||||
r.Responses[ip] = resp
|
||||
order = append(order, ip)
|
||||
if resp.Status == StatusRestart || resp.Status == StatusReferral {
|
||||
childsets++
|
||||
}
|
||||
}
|
||||
|
||||
// Phase two: create the children.
|
||||
childset := 0
|
||||
var sets [][]*Referral
|
||||
for _, ip := range order {
|
||||
resp := r.Responses[ip]
|
||||
if resp.Status != StatusRestart && resp.Status != StatusReferral {
|
||||
continue
|
||||
}
|
||||
childset++
|
||||
refid := r.RefID
|
||||
if childsets > 1 {
|
||||
refid = fmt.Sprintf("%s.%d", r.RefID, childset)
|
||||
}
|
||||
children := r.makeReferrals(resp, refid, ip)
|
||||
r.Children[ip] = children
|
||||
sets = append(sets, children)
|
||||
}
|
||||
return sets, nil
|
||||
}
|
||||
|
||||
// processAddRoots creates one child per root server with equal weight
|
||||
// (referral.rb process_add_roots); the roots come from the info cache hints.
|
||||
func (r *Referral) processAddRoots() ([]*Referral, error) {
|
||||
starters, _, err := r.InfoCache.GetStartServers("")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
dot := ""
|
||||
if r.RefID != "" {
|
||||
dot = "."
|
||||
}
|
||||
var children []*Referral
|
||||
for i, root := range starters {
|
||||
child := r.makeReferral(referralArgs{
|
||||
server: root.Name,
|
||||
serverIPs: root.IPs,
|
||||
refid: fmt.Sprintf("%s%s%d", r.RefID, dot, i+1),
|
||||
})
|
||||
children = append(children, child)
|
||||
}
|
||||
r.Children["rootroot"] = children
|
||||
return children, nil
|
||||
}
|
||||
|
||||
// makeReferrals creates one child per start server for a referral/restart
|
||||
// response (referral.rb make_referrals): qname moves to the response's
|
||||
// endname (the CNAME target on restart), the bailiwick and cache come from
|
||||
// the response.
|
||||
func (r *Referral) makeReferrals(resp *ServerResponse, refid, parentIP string) []*Referral {
|
||||
var children []*Referral
|
||||
for i, st := range resp.Starters {
|
||||
children = append(children, r.makeReferral(referralArgs{
|
||||
qname: resp.DQ.Endname,
|
||||
server: st.Name,
|
||||
serverIPs: st.IPs,
|
||||
bailiwick: resp.StartersBailiwick,
|
||||
infoCache: resp.Cache,
|
||||
refid: fmt.Sprintf("%s.%d", refid, i+1),
|
||||
parentIP: parentIP,
|
||||
}))
|
||||
}
|
||||
return children
|
||||
}
|
||||
|
||||
// replaceChild swaps before for after in the children/resolves lists (fast
|
||||
// mode substitution); before keeps a pointer to its replacement.
|
||||
func (r *Referral) replaceChild(before, after *Referral) {
|
||||
before.ReplacedBy = after
|
||||
for ip := range r.Children {
|
||||
for i, c := range r.Children[ip] {
|
||||
if c == before {
|
||||
r.Children[ip][i] = after
|
||||
}
|
||||
}
|
||||
}
|
||||
for i, c := range r.Resolves {
|
||||
if c == before {
|
||||
r.Resolves[i] = after
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+176
-141
@@ -1,159 +1,194 @@
|
||||
package traverse
|
||||
|
||||
import (
|
||||
"net"
|
||||
"testing"
|
||||
|
||||
"gitea.hansenits.com.au/hits/ExploreDNS/internal/dns"
|
||||
"github.com/miekg/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.")
|
||||
func newTestReferral(server string, ips []string) *Referral {
|
||||
r := &Referral{
|
||||
RefID: "1",
|
||||
Qname: "www.example.com",
|
||||
Qclass: dns.ClassINET,
|
||||
Qtype: dns.TypeA,
|
||||
NSAType: dns.TypeA,
|
||||
Server: server,
|
||||
ServerIPs: ips,
|
||||
Bailiwick: "com",
|
||||
InfoCache: NewInfoCache(nil),
|
||||
Status: RefStatusNormal,
|
||||
Responses: make(map[string]*ServerResponse),
|
||||
Children: make(map[string][]*Referral),
|
||||
ServerWeights: make(map[string]float64),
|
||||
maxdepth: DefaultMaxDepth,
|
||||
}
|
||||
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")
|
||||
for _, ip := range ips {
|
||||
r.ServerWeights[ip] = 1.0 / float64(len(ips))
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
func TestReferralInBailiwick(t *testing.T) {
|
||||
func TestRefidDepth(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
bailiwick string
|
||||
testName string
|
||||
want bool
|
||||
refid string
|
||||
want int
|
||||
}{
|
||||
{"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},
|
||||
{"", 0},
|
||||
{"1", 1},
|
||||
{"1.1.2", 3},
|
||||
{"1.2.0.1", 3},
|
||||
{"1.1.2.0.1.4.2.0.2.0.2", 8},
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCircularReferralError(t *testing.T) {
|
||||
err := &CircularReferralError{
|
||||
Name: "ns.example.com.",
|
||||
Chain: []string{"ns1.example.com.", "ns2.example.com."},
|
||||
}
|
||||
|
||||
expected := "circular referral detected for ns.example.com.: [ns1.example.com. ns2.example.com.]"
|
||||
if err.Error() != expected {
|
||||
t.Errorf("Error() = %q, want %q", err.Error(), expected)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUnresolvableNameserverError(t *testing.T) {
|
||||
err := &UnresolvableNameserverError{
|
||||
Name: "ns.example.com.",
|
||||
Reason: "NXDOMAIN",
|
||||
}
|
||||
|
||||
expected := "unresolvable nameserver ns.example.com.: NXDOMAIN"
|
||||
if err.Error() != expected {
|
||||
t.Errorf("Error() = %q, want %q", err.Error(), expected)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetVisitedNames(t *testing.T) {
|
||||
visited := map[string]bool{
|
||||
"ns1.example.com.": true,
|
||||
"ns2.example.com.": true,
|
||||
"ns3.example.com.": true,
|
||||
}
|
||||
|
||||
names := getVisitedNames(visited)
|
||||
if len(names) != 3 {
|
||||
t.Errorf("got %d names, want 3", len(names))
|
||||
}
|
||||
|
||||
seen := make(map[string]bool)
|
||||
for _, name := range names {
|
||||
if seen[name] {
|
||||
t.Errorf("duplicate name: %s", name)
|
||||
}
|
||||
seen[name] = true
|
||||
if !visited[name] {
|
||||
t.Errorf("unexpected name: %s", name)
|
||||
if got := refidDepth(tt.refid); got != tt.want {
|
||||
t.Errorf("refidDepth(%q) = %d, want %d", tt.refid, got, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestTxtIPsVerbose(t *testing.T) {
|
||||
r := newTestReferral("ns1.example.com", []string{"2.2.2.2", "1.1.1.1"})
|
||||
if got := r.TxtIPsVerbose(); got != "50.0%=1.1.1.1,50.0%=2.2.2.2" {
|
||||
t.Errorf("TxtIPsVerbose = %q", got)
|
||||
}
|
||||
|
||||
// key: pseudo entries keep only the first two fields.
|
||||
r2 := newTestReferral("ns2.example.com", nil)
|
||||
r2.ServerIPs = []string{"key:noglue:9.9.9.9:www.example.com:IN:A:x:y"}
|
||||
r2.ServerWeights = map[string]float64{r2.ServerIPs[0]: 1.0}
|
||||
if got := r2.TxtIPsVerbose(); got != "100.0%=noglue:9.9.9.9" {
|
||||
t.Errorf("TxtIPsVerbose key entry = %q", got)
|
||||
}
|
||||
|
||||
var unresolved Referral
|
||||
if got := unresolved.TxtIPsVerbose(); got != "" {
|
||||
t.Errorf("unresolved TxtIPsVerbose = %q, want empty", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFastKeyLowercases(t *testing.T) {
|
||||
r := newTestReferral("NS1.Example.COM", []string{"1.1.1.1"})
|
||||
r.Server = "NS1.Example.COM" // bypass canonicalisation to prove downcasing
|
||||
key := fastKey(r)
|
||||
if key != "www.example.com:in:a:ns1.example.com:100.0%=1.1.1.1" {
|
||||
t.Errorf("fastKey = %q", key)
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsNoGlueAndIsLoop(t *testing.T) {
|
||||
inBailiwick := newTestReferral("ns1.sub.com", nil)
|
||||
if !inBailiwick.isNoGlue() {
|
||||
t.Error("in-bailiwick NS without addresses should be noglue")
|
||||
}
|
||||
outOfBailiwick := newTestReferral("ns1.other.net", nil)
|
||||
if outOfBailiwick.isNoGlue() {
|
||||
t.Error("out-of-bailiwick NS is resolvable, not noglue")
|
||||
}
|
||||
resolved := newTestReferral("ns1.sub.com", []string{"1.1.1.1"})
|
||||
if resolved.isNoGlue() || resolved.isLoop() {
|
||||
t.Error("resolved referral is neither noglue nor loop")
|
||||
}
|
||||
|
||||
parent := newTestReferral("ns.b.net", nil)
|
||||
parent.Qname = "ns.a.net"
|
||||
child := newTestReferral("ns.b.net", nil)
|
||||
child.Qname = "ns.a.net"
|
||||
child.Parent = parent
|
||||
if !child.isLoop() {
|
||||
t.Error("same qname/qclass/qtype/server with unresolved ancestor should loop")
|
||||
}
|
||||
child.Server = "ns.c.net"
|
||||
if child.isLoop() {
|
||||
t.Error("different server must not loop")
|
||||
}
|
||||
}
|
||||
|
||||
func TestChainHasQuery(t *testing.T) {
|
||||
parent := newTestReferral("ns1.example.com", []string{"1.1.1.1"})
|
||||
parent.Qname = "www.a.com"
|
||||
child := newTestReferral("ns2.example.com", []string{"2.2.2.2"})
|
||||
child.Qname = "www.b.net"
|
||||
child.Parent = parent
|
||||
if !child.chainHasQuery("www.a.com") {
|
||||
t.Error("ancestor qname should be found")
|
||||
}
|
||||
if !child.chainHasQuery("WWW.B.NET.") {
|
||||
t.Error("own qname should be found case-insensitively")
|
||||
}
|
||||
if child.chainHasQuery("www.c.org") {
|
||||
t.Error("unknown qname must not match")
|
||||
}
|
||||
}
|
||||
|
||||
func TestReplaceChild(t *testing.T) {
|
||||
parent := newTestReferral("ns1.example.com", []string{"1.1.1.1"})
|
||||
before := newTestReferral("ns2.example.com", []string{"2.2.2.2"})
|
||||
after := newTestReferral("ns2.example.com", []string{"2.2.2.2"})
|
||||
parent.Children["1.1.1.1"] = []*Referral{before}
|
||||
parent.Resolves = []*Referral{before}
|
||||
|
||||
parent.replaceChild(before, after)
|
||||
if parent.Children["1.1.1.1"][0] != after || parent.Resolves[0] != after {
|
||||
t.Error("replaceChild did not swap the node everywhere")
|
||||
}
|
||||
if before.ReplacedBy != after {
|
||||
t.Error("replaced node should point at its replacement")
|
||||
}
|
||||
}
|
||||
|
||||
func TestOverallStatus(t *testing.T) {
|
||||
r := newTestReferral("ns1.example.com", nil)
|
||||
r.Status = RefStatusNoGlue
|
||||
if got := r.OverallStatus(); got != StatusNoGlue {
|
||||
t.Errorf("noglue overall = %q", got)
|
||||
}
|
||||
r.Status = RefStatusLoop
|
||||
if got := r.OverallStatus(); got != StatusLoop {
|
||||
t.Errorf("loop overall = %q", got)
|
||||
}
|
||||
|
||||
n := newTestReferral("ns1.example.com", []string{"1.1.1.1", "2.2.2.2"})
|
||||
if got := n.OverallStatus(); got != "" {
|
||||
t.Errorf("no responses overall = %q, want empty", got)
|
||||
}
|
||||
n.Responses["1.1.1.1"] = &ServerResponse{Status: StatusAnswered}
|
||||
if got := n.OverallStatus(); got != StatusAnswered {
|
||||
t.Errorf("single status overall = %q", got)
|
||||
}
|
||||
n.Responses["2.2.2.2"] = &ServerResponse{Status: StatusError}
|
||||
if got := n.OverallStatus(); got != "mixed" {
|
||||
t.Errorf("mixed overall = %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestIPsAsArraySkipsPseudoKeys(t *testing.T) {
|
||||
r := newTestReferral("ns1.example.com", []string{"1.1.1.1", "key:noglue:2.2.2.2"})
|
||||
got := r.IPsAsArray()
|
||||
if len(got) != 1 || got[0] != "1.1.1.1" {
|
||||
t.Errorf("IPsAsArray = %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestToASCII(t *testing.T) {
|
||||
if got := toASCII("bücher.example"); got != "xn--bcher-kva.example" {
|
||||
t.Errorf("toASCII = %q", got)
|
||||
}
|
||||
if got := toASCII("plain.example"); got != "plain.example" {
|
||||
t.Errorf("ascii name changed: %q", got)
|
||||
}
|
||||
if got := toASCII(""); got != "" {
|
||||
t.Errorf("empty name changed: %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestServerResponseString(t *testing.T) {
|
||||
noglue := NewNoGlueResponse("www.example.com", dns.ClassINET, dns.TypeA, "1.1.1.1", "ns1.example.com", "example.com")
|
||||
if got := noglue.String(); got != "No glue for ns1.example.com" {
|
||||
t.Errorf("noglue String = %q", got)
|
||||
}
|
||||
loop := NewLoopResponse("www.example.com", dns.ClassINET, dns.TypeA, "1.1.1.1", "ns1.example.com", "example.com")
|
||||
if got := loop.String(); got != "Loop encountered resolving ns1.example.com" {
|
||||
t.Errorf("loop String = %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,256 +0,0 @@
|
||||
package traverse
|
||||
|
||||
import (
|
||||
"net"
|
||||
|
||||
"gitea.hansenits.com.au/hits/ExploreDNS/internal/dns"
|
||||
miekgdns "github.com/miekg/dns"
|
||||
)
|
||||
|
||||
type ResponseType int
|
||||
|
||||
const (
|
||||
RespReferral ResponseType = iota
|
||||
RespAnswer
|
||||
RespCNAMEFollow
|
||||
RespNODATA
|
||||
RespNXDOMAIN
|
||||
RespSERVFAIL
|
||||
RespREFUSED
|
||||
RespNOTIMPL
|
||||
RespCNAMELoop
|
||||
RespError
|
||||
// RespNSResolutionFailed indicates that the traversal could not resolve the
|
||||
// IP address of an in-bailiwick nameserver. The domain may still be
|
||||
// reachable in practice (e.g. via glue records held by the registry), but
|
||||
// the iterative traversal could not complete that path.
|
||||
RespNSResolutionFailed
|
||||
)
|
||||
|
||||
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 RespREFUSED:
|
||||
return "refused"
|
||||
case RespNOTIMPL:
|
||||
return "notimp"
|
||||
case RespCNAMELoop:
|
||||
return "cname_loop"
|
||||
case RespError:
|
||||
return "error"
|
||||
case RespNSResolutionFailed:
|
||||
return "ns_error"
|
||||
default:
|
||||
return "unknown"
|
||||
}
|
||||
}
|
||||
|
||||
type Response struct {
|
||||
Referral *Referral
|
||||
Server net.IP
|
||||
Cache *InfoCache
|
||||
Decoded *dns.DecodedResponse
|
||||
Type ResponseType
|
||||
ErrorMessage string
|
||||
}
|
||||
|
||||
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
|
||||
r.ErrorMessage = "nil DNS response"
|
||||
return r
|
||||
}
|
||||
|
||||
r.Decoded = dns.DecodeResponse(msg)
|
||||
if r.Decoded == nil {
|
||||
r.Type = RespError
|
||||
r.ErrorMessage = "failed to decode DNS response"
|
||||
return r
|
||||
}
|
||||
|
||||
// Synthesize CNAME from DNAME when the server didn't include a synthesized CNAME record.
|
||||
if len(r.Decoded.CNAMEChain) == 0 && r.Referral != nil && len(r.Decoded.DNAMEMappings) > 0 {
|
||||
for _, dm := range r.Decoded.DNAMEMappings {
|
||||
synthesized := dns.SynthesizeCNAMEFromDNAME(r.Referral.Name, dm.Owner, dm.Target)
|
||||
if synthesized != "" {
|
||||
r.Decoded.CNAMEChain = append(r.Decoded.CNAMEChain, synthesized)
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
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.ResponseREFUSED:
|
||||
return RespREFUSED
|
||||
case dns.ResponseNOTIMPL:
|
||||
return RespNOTIMPL
|
||||
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 {
|
||||
switch rr.(type) {
|
||||
case *miekgdns.CNAME, *miekgdns.DNAME, *miekgdns.RRSIG:
|
||||
// CNAME and DNAME are redirect records, not final answers.
|
||||
// RRSIG is a DNSSEC signature record — it covers the CNAME/DNAME
|
||||
// but is not itself the answer to the original question type.
|
||||
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, RespREFUSED, RespNOTIMPL, RespCNAMELoop, RespError, RespNSResolutionFailed:
|
||||
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)
|
||||
}
|
||||
@@ -1,340 +0,0 @@
|
||||
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},
|
||||
{RespNSResolutionFailed, 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"},
|
||||
{RespNSResolutionFailed, "ns_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")
|
||||
}
|
||||
}
|
||||
@@ -1,768 +0,0 @@
|
||||
package traverse
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
|
||||
"github.com/miekg/dns"
|
||||
)
|
||||
|
||||
// TestCNAMELoopDetected verifies that a two-step CNAME loop (A → B → A) is
|
||||
// detected without infinite recursion and produces a RespCNAMELoop result.
|
||||
func TestCNAMELoopDetected(t *testing.T) {
|
||||
// www.example.com → CNAME → alias.example.com → CNAME → www.example.com (loop)
|
||||
cnameToAlias := new(dns.Msg)
|
||||
cnameToAlias.SetReply(new(dns.Msg))
|
||||
cnameToAlias.Answer = append(cnameToAlias.Answer, &dns.CNAME{
|
||||
Hdr: dns.RR_Header{Name: "www.example.com.", Rrtype: dns.TypeCNAME, Class: dns.ClassINET},
|
||||
Target: "alias.example.com.",
|
||||
})
|
||||
|
||||
cnameBack := new(dns.Msg)
|
||||
cnameBack.SetReply(new(dns.Msg))
|
||||
cnameBack.Answer = append(cnameBack.Answer, &dns.CNAME{
|
||||
Hdr: dns.RR_Header{Name: "alias.example.com.", Rrtype: dns.TypeCNAME, Class: dns.ClassINET},
|
||||
Target: "www.example.com.",
|
||||
})
|
||||
|
||||
tr := NewTraverser(&TraverserConfig{
|
||||
MaxDepth: 10,
|
||||
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]
|
||||
switch q.Name {
|
||||
case "www.example.com.":
|
||||
return cnameToAlias.Copy(), nil
|
||||
case "alias.example.com.":
|
||||
return cnameBack.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)
|
||||
}
|
||||
|
||||
foundLoop := false
|
||||
for _, r := range results {
|
||||
if r.Response != nil && r.Response.Type == RespCNAMELoop {
|
||||
foundLoop = true
|
||||
if r.Response.ErrorMessage == "" {
|
||||
t.Error("expected non-empty ErrorMessage on CNAME loop result")
|
||||
}
|
||||
}
|
||||
}
|
||||
if !foundLoop {
|
||||
t.Error("expected RespCNAMELoop result for CNAME loop A → B → A")
|
||||
}
|
||||
}
|
||||
|
||||
// TestCNAMEDirectLoop verifies that a direct self-loop (A → A) is handled.
|
||||
func TestCNAMEDirectLoop(t *testing.T) {
|
||||
selfLoop := new(dns.Msg)
|
||||
selfLoop.SetReply(new(dns.Msg))
|
||||
selfLoop.Answer = append(selfLoop.Answer, &dns.CNAME{
|
||||
Hdr: dns.RR_Header{Name: "www.example.com.", Rrtype: dns.TypeCNAME, Class: dns.ClassINET},
|
||||
Target: "www.example.com.",
|
||||
})
|
||||
|
||||
tr := NewTraverser(&TraverserConfig{
|
||||
MaxDepth: 10,
|
||||
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 selfLoop.Copy(), nil
|
||||
})
|
||||
|
||||
ctx := context.Background()
|
||||
results, err := tr.Traverse(ctx, "www.example.com")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
|
||||
foundLoop := false
|
||||
for _, r := range results {
|
||||
if r.Response != nil && r.Response.Type == RespCNAMELoop {
|
||||
foundLoop = true
|
||||
}
|
||||
}
|
||||
if !foundLoop {
|
||||
t.Error("expected RespCNAMELoop for direct self-referencing CNAME")
|
||||
}
|
||||
}
|
||||
|
||||
// TestREFUSEDResponse verifies that a REFUSED rcode is classified as RespREFUSED.
|
||||
func TestREFUSEDResponse(t *testing.T) {
|
||||
refusedResp := new(dns.Msg)
|
||||
refusedResp.Rcode = dns.RcodeRefused
|
||||
|
||||
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 refusedResp.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 != RespREFUSED {
|
||||
t.Errorf("Type = %s, want refused", results[0].Response.Type)
|
||||
}
|
||||
if !results[0].Response.IsTerminal() {
|
||||
t.Error("REFUSED should be a terminal response")
|
||||
}
|
||||
}
|
||||
|
||||
// TestNOTIMPLResponse verifies that a NOTIMP rcode is classified as RespNOTIMPL.
|
||||
func TestNOTIMPLResponse(t *testing.T) {
|
||||
notImplResp := new(dns.Msg)
|
||||
notImplResp.Rcode = dns.RcodeNotImplemented
|
||||
|
||||
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 notImplResp.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 != RespNOTIMPL {
|
||||
t.Errorf("Type = %s, want notimp", results[0].Response.Type)
|
||||
}
|
||||
if !results[0].Response.IsTerminal() {
|
||||
t.Error("NOTIMP should be a terminal response")
|
||||
}
|
||||
}
|
||||
|
||||
// TestGracefulDegradationUnreachableServer verifies that when some servers are
|
||||
// unreachable, traversal continues with the remaining servers and does not panic.
|
||||
func TestGracefulDegradationUnreachableServer(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"),
|
||||
})
|
||||
|
||||
// Referral with two nameservers; first always fails, second provides the answer.
|
||||
referralMsg := new(dns.Msg)
|
||||
referralMsg.Rcode = dns.RcodeSuccess
|
||||
referralMsg.Authoritative = false
|
||||
referralMsg.Ns = append(referralMsg.Ns,
|
||||
&dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns1.example.com."},
|
||||
&dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns2.example.com."},
|
||||
)
|
||||
referralMsg.Extra = append(referralMsg.Extra,
|
||||
&dns.A{Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.1")},
|
||||
&dns.A{Hdr: dns.RR_Header{Name: "ns2.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.2")},
|
||||
)
|
||||
|
||||
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) {
|
||||
switch server {
|
||||
case "198.41.0.4":
|
||||
return referralMsg.Copy(), nil
|
||||
case "10.0.0.1":
|
||||
return nil, errors.New("connection refused")
|
||||
case "10.0.0.2":
|
||||
return answerResp.Copy(), nil
|
||||
}
|
||||
return nil, errors.New("unexpected server")
|
||||
})
|
||||
|
||||
ctx := context.Background()
|
||||
results, err := tr.Traverse(ctx, "example.com")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
|
||||
foundAnswer := false
|
||||
for _, r := range results {
|
||||
if r.Response != nil && r.Response.Type == RespAnswer {
|
||||
foundAnswer = true
|
||||
}
|
||||
}
|
||||
if !foundAnswer {
|
||||
t.Error("expected an answer result from the reachable server")
|
||||
}
|
||||
}
|
||||
|
||||
// TestGracefulDegradationAllUnreachable verifies that when ALL servers fail,
|
||||
// the traversal returns a SERVFAIL result without panicking.
|
||||
func TestGracefulDegradationAllUnreachable(t *testing.T) {
|
||||
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 nil, errors.New("network unreachable")
|
||||
})
|
||||
|
||||
ctx := context.Background()
|
||||
results, err := tr.Traverse(ctx, "example.com")
|
||||
if err != nil {
|
||||
t.Fatalf("traversal must not return a top-level error: %v", err)
|
||||
}
|
||||
if len(results) == 0 {
|
||||
t.Fatal("expected at least one result even on total failure")
|
||||
}
|
||||
last := results[len(results)-1]
|
||||
if last.Response == nil {
|
||||
t.Fatal("last result must have a response")
|
||||
}
|
||||
if last.Response.Type != RespSERVFAIL && last.Response.Type != RespError {
|
||||
t.Errorf("expected SERVFAIL or error when all servers unreachable, got %s", last.Response.Type)
|
||||
}
|
||||
}
|
||||
|
||||
// TestDNAMEFollowNoSynthesizedCNAME verifies that a DNAME record in the answer
|
||||
// section synthesizes a CNAME follow when the server doesn't include one.
|
||||
func TestDNAMEFollowNoSynthesizedCNAME(t *testing.T) {
|
||||
// Server returns DNAME only (no synthesized CNAME).
|
||||
dnameResp := new(dns.Msg)
|
||||
dnameResp.SetReply(new(dns.Msg))
|
||||
dnameResp.Answer = append(dnameResp.Answer, &dns.DNAME{
|
||||
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeDNAME, Class: dns.ClassINET, Ttl: 300},
|
||||
Target: "example.net.",
|
||||
})
|
||||
|
||||
answerResp := new(dns.Msg)
|
||||
answerResp.SetReply(new(dns.Msg))
|
||||
answerResp.Answer = append(answerResp.Answer, &dns.A{
|
||||
Hdr: dns.RR_Header{Name: "www.example.net.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300},
|
||||
A: net.ParseIP("203.0.113.1"),
|
||||
})
|
||||
|
||||
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 dnameResp.Copy(), nil
|
||||
}
|
||||
if q.Name == "www.example.net." {
|
||||
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)
|
||||
}
|
||||
|
||||
foundCNAMEFollow := false
|
||||
for _, r := range results {
|
||||
if r.Response != nil && r.Response.Type == RespCNAMEFollow {
|
||||
foundCNAMEFollow = true
|
||||
}
|
||||
}
|
||||
if !foundCNAMEFollow {
|
||||
t.Error("expected RespCNAMEFollow synthesized from DNAME record")
|
||||
}
|
||||
}
|
||||
|
||||
// TestIsNameInChain verifies the ancestor chain lookup.
|
||||
func TestIsNameInChain(t *testing.T) {
|
||||
root := NewReferral("example.com", dnsTypeA, ".", 0, 1.0, nil)
|
||||
child := NewReferral("www.example.com", dnsTypeA, "example.com.", 1, 1.0, root)
|
||||
grandchild := NewReferral("sub.www.example.com", dnsTypeA, "www.example.com.", 2, 1.0, child)
|
||||
|
||||
tests := []struct {
|
||||
ref *Referral
|
||||
name string
|
||||
want bool
|
||||
}{
|
||||
{grandchild, "sub.www.example.com", true}, // self
|
||||
{grandchild, "www.example.com", true}, // parent
|
||||
{grandchild, "example.com", true}, // grandparent
|
||||
{grandchild, "other.example.com", false}, // not in chain
|
||||
{root, "example.com", true}, // root matches itself
|
||||
{root, "www.example.com", false}, // child not in chain from root
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
got := tt.ref.IsNameInChain(tt.name)
|
||||
if got != tt.want {
|
||||
t.Errorf("IsNameInChain(%q) from %q = %v, want %v", tt.name, tt.ref.Name, got, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestResponseTypeStrings verifies String() for new response types.
|
||||
func TestResponseTypeStrings(t *testing.T) {
|
||||
tests := []struct {
|
||||
rt ResponseType
|
||||
want string
|
||||
}{
|
||||
{RespReferral, "referral"},
|
||||
{RespAnswer, "answer"},
|
||||
{RespCNAMEFollow, "cname_follow"},
|
||||
{RespNODATA, "nodata"},
|
||||
{RespNXDOMAIN, "nxdomain"},
|
||||
{RespSERVFAIL, "servfail"},
|
||||
{RespREFUSED, "refused"},
|
||||
{RespNOTIMPL, "notimp"},
|
||||
{RespCNAMELoop, "cname_loop"},
|
||||
{RespError, "error"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
if got := tt.rt.String(); got != tt.want {
|
||||
t.Errorf("ResponseType(%d).String() = %q, want %q", tt.rt, got, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestMalformedResponseNoPanic verifies that a nil response from the exchange
|
||||
// function does not cause a panic, and produces an error result.
|
||||
func TestMalformedResponseNoPanic(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 // nil response, no error
|
||||
})
|
||||
|
||||
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 one result")
|
||||
}
|
||||
// Should produce error/servfail, not panic
|
||||
for _, r := range results {
|
||||
if r.Response == nil {
|
||||
t.Error("result has nil response")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestDNSSECRRSIGDoesNotBlockCNAMEFollow verifies that a DNSSEC RRSIG record
|
||||
// accompanying a CNAME in the answer section is treated as metadata and does
|
||||
// NOT prevent the traversal from following the CNAME.
|
||||
func TestDNSSECRRSIGDoesNotBlockCNAMEFollow(t *testing.T) {
|
||||
// Server returns CNAME + RRSIG (DNSSEC-signed zone response).
|
||||
cnameWithRRSIG := new(dns.Msg)
|
||||
cnameWithRRSIG.SetReply(new(dns.Msg))
|
||||
cnameWithRRSIG.Answer = append(cnameWithRRSIG.Answer,
|
||||
&dns.CNAME{
|
||||
Hdr: dns.RR_Header{Name: "www.example.com.", Rrtype: dns.TypeCNAME, Class: dns.ClassINET, Ttl: 300},
|
||||
Target: "example.com.",
|
||||
},
|
||||
&dns.RRSIG{
|
||||
Hdr: dns.RR_Header{Name: "www.example.com.", Rrtype: dns.TypeRRSIG, Class: dns.ClassINET, Ttl: 300},
|
||||
TypeCovered: dns.TypeCNAME,
|
||||
},
|
||||
)
|
||||
|
||||
finalAnswer := new(dns.Msg)
|
||||
finalAnswer.SetReply(new(dns.Msg))
|
||||
finalAnswer.Answer = append(finalAnswer.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: 10,
|
||||
QueryType: dnsTypeA,
|
||||
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||
Fast: true,
|
||||
})
|
||||
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 cnameWithRRSIG.Copy(), nil
|
||||
}
|
||||
return finalAnswer.Copy(), nil
|
||||
})
|
||||
|
||||
ctx := context.Background()
|
||||
results, err := tr.Traverse(ctx, "www.example.com")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
|
||||
foundCNAMEFollow := false
|
||||
foundAnswer := false
|
||||
for _, r := range results {
|
||||
if r.Response != nil {
|
||||
switch r.Response.Type {
|
||||
case RespCNAMEFollow:
|
||||
foundCNAMEFollow = true
|
||||
case RespAnswer:
|
||||
foundAnswer = true
|
||||
}
|
||||
}
|
||||
}
|
||||
if !foundCNAMEFollow {
|
||||
t.Error("expected RespCNAMEFollow: RRSIG should not block CNAME following")
|
||||
}
|
||||
if !foundAnswer {
|
||||
t.Error("expected final RespAnswer after CNAME follow")
|
||||
}
|
||||
}
|
||||
|
||||
// TestFastModeOn verifies that Fast=true uses the shared root cache (default
|
||||
// behaviour): a child branch can see glue stored by the root referral.
|
||||
func TestFastModeOn(t *testing.T) {
|
||||
// Root referral returns two nameservers with glue. Each NS branch returns
|
||||
// an answer. We verify both branches are queried.
|
||||
referralMsg := new(dns.Msg)
|
||||
referralMsg.Rcode = dns.RcodeSuccess
|
||||
referralMsg.Authoritative = false
|
||||
referralMsg.Ns = append(referralMsg.Ns,
|
||||
&dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns1.example.com."},
|
||||
&dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns2.example.com."},
|
||||
)
|
||||
referralMsg.Extra = append(referralMsg.Extra,
|
||||
&dns.A{Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.1")},
|
||||
&dns.A{Hdr: dns.RR_Header{Name: "ns2.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.2")},
|
||||
)
|
||||
|
||||
answerMsg := new(dns.Msg)
|
||||
answerMsg.SetReply(new(dns.Msg))
|
||||
answerMsg.Answer = append(answerMsg.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")},
|
||||
Fast: true,
|
||||
})
|
||||
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||
if server == "198.41.0.4" {
|
||||
return referralMsg.Copy(), nil
|
||||
}
|
||||
return answerMsg.Copy(), nil
|
||||
})
|
||||
|
||||
ctx := context.Background()
|
||||
results, err := tr.Traverse(ctx, "example.com")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
|
||||
answers := 0
|
||||
for _, r := range results {
|
||||
if r.Response != nil && r.Response.Type == RespAnswer {
|
||||
answers++
|
||||
}
|
||||
}
|
||||
if answers == 0 {
|
||||
t.Error("expected at least one answer with Fast=true")
|
||||
}
|
||||
}
|
||||
|
||||
// TestFastModeOff verifies that Fast=false gives each referral its own
|
||||
// independent cache — no cross-branch glue contamination.
|
||||
func TestFastModeOff(t *testing.T) {
|
||||
referralMsg := new(dns.Msg)
|
||||
referralMsg.Rcode = dns.RcodeSuccess
|
||||
referralMsg.Authoritative = false
|
||||
referralMsg.Ns = append(referralMsg.Ns,
|
||||
&dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns1.example.com."},
|
||||
&dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns2.example.com."},
|
||||
)
|
||||
referralMsg.Extra = append(referralMsg.Extra,
|
||||
&dns.A{Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.1")},
|
||||
&dns.A{Hdr: dns.RR_Header{Name: "ns2.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.2")},
|
||||
)
|
||||
|
||||
answerMsg := new(dns.Msg)
|
||||
answerMsg.SetReply(new(dns.Msg))
|
||||
answerMsg.Answer = append(answerMsg.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")},
|
||||
Fast: false,
|
||||
})
|
||||
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||
if server == "198.41.0.4" {
|
||||
return referralMsg.Copy(), nil
|
||||
}
|
||||
return answerMsg.Copy(), nil
|
||||
})
|
||||
|
||||
ctx := context.Background()
|
||||
results, err := tr.Traverse(ctx, "example.com")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
// Traversal must complete without panic and produce results.
|
||||
if len(results) == 0 {
|
||||
t.Fatal("expected at least one result with Fast=false")
|
||||
}
|
||||
}
|
||||
|
||||
// TestFastModeDefaultIsTrue verifies that DefaultTraverserConfig has Fast=true.
|
||||
func TestFastModeDefaultIsTrue(t *testing.T) {
|
||||
cfg := DefaultTraverserConfig()
|
||||
if !cfg.Fast {
|
||||
t.Error("DefaultTraverserConfig().Fast should be true")
|
||||
}
|
||||
}
|
||||
|
||||
// TestManyNSRecords verifies that a referral with more than 10 nameservers is
|
||||
// handled gracefully — no panics, results are produced.
|
||||
func TestManyNSRecords(t *testing.T) {
|
||||
referralMsg := new(dns.Msg)
|
||||
referralMsg.Rcode = dns.RcodeSuccess
|
||||
referralMsg.Authoritative = false
|
||||
for i := 1; i <= 12; i++ {
|
||||
ns := &dns.NS{
|
||||
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS},
|
||||
Ns: net.ParseIP(string(rune('a'+i-1))).String() + ".ns.example.com.",
|
||||
}
|
||||
// Use a distinct IP for each NS so glue is resolved.
|
||||
ip := net.IP{10, 0, 0, byte(i)}
|
||||
referralMsg.Ns = append(referralMsg.Ns, &dns.NS{
|
||||
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS},
|
||||
Ns: ns.Ns,
|
||||
})
|
||||
referralMsg.Extra = append(referralMsg.Extra, &dns.A{
|
||||
Hdr: dns.RR_Header{Name: ns.Ns, Rrtype: dnsTypeA},
|
||||
A: ip,
|
||||
})
|
||||
}
|
||||
|
||||
answerMsg := new(dns.Msg)
|
||||
answerMsg.SetReply(new(dns.Msg))
|
||||
answerMsg.Answer = append(answerMsg.Answer, &dns.A{
|
||||
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300},
|
||||
A: net.ParseIP("93.184.216.34"),
|
||||
})
|
||||
|
||||
var queries int64
|
||||
tr := NewTraverser(&TraverserConfig{
|
||||
MaxDepth: 5,
|
||||
QueryType: dnsTypeA,
|
||||
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||
Fast: true,
|
||||
})
|
||||
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||
atomic.AddInt64(&queries, 1)
|
||||
if server == "198.41.0.4" {
|
||||
return referralMsg.Copy(), nil
|
||||
}
|
||||
return answerMsg.Copy(), nil
|
||||
})
|
||||
|
||||
ctx := context.Background()
|
||||
results, err := tr.Traverse(ctx, "example.com")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error with 12 NS records: %v", err)
|
||||
}
|
||||
if len(results) == 0 {
|
||||
t.Fatal("expected results with many NS records")
|
||||
}
|
||||
foundAnswer := false
|
||||
for _, r := range results {
|
||||
if r.Response != nil && r.Response.Type == RespAnswer {
|
||||
foundAnswer = true
|
||||
}
|
||||
}
|
||||
if !foundAnswer {
|
||||
t.Error("expected at least one answer from the 12-NS referral")
|
||||
}
|
||||
}
|
||||
|
||||
// TestIDNPunycodeConversion verifies that a unicode (IDN) domain name is
|
||||
// converted to its punycode/ACE form before querying.
|
||||
func TestIDNPunycodeConversion(t *testing.T) {
|
||||
// "münchen.de" → "xn--mnchen-3ya.de" (after punycode encoding)
|
||||
ref := NewReferral("münchen.de", dnsTypeA, ".", 0, 1.0, nil)
|
||||
if ref.Name == "münchen.de." {
|
||||
t.Errorf("IDN name was not converted to punycode: got %q", ref.Name)
|
||||
}
|
||||
// Verify it starts with the expected punycode label.
|
||||
if ref.Name != "xn--mnchen-3ya.de." {
|
||||
t.Errorf("unexpected punycode result: got %q, want %q", ref.Name, "xn--mnchen-3ya.de.")
|
||||
}
|
||||
}
|
||||
|
||||
// TestASCIIDomainUnchanged verifies that a plain ASCII domain is not mangled
|
||||
// by the IDN conversion path.
|
||||
func TestASCIIDomainUnchanged(t *testing.T) {
|
||||
ref := NewReferral("example.com", dnsTypeA, ".", 0, 1.0, nil)
|
||||
if ref.Name != "example.com." {
|
||||
t.Errorf("ASCII domain was mangled: got %q, want %q", ref.Name, "example.com.")
|
||||
}
|
||||
}
|
||||
|
||||
// TestWildcardResponse verifies that a wildcard answer (e.g. *.example.com
|
||||
// returning an A record for sub.example.com) is handled as a regular answer.
|
||||
func TestWildcardResponse(t *testing.T) {
|
||||
wildcardAnswer := new(dns.Msg)
|
||||
wildcardAnswer.SetReply(new(dns.Msg))
|
||||
wildcardAnswer.Authoritative = true
|
||||
wildcardAnswer.Answer = append(wildcardAnswer.Answer, &dns.A{
|
||||
Hdr: dns.RR_Header{Name: "sub.example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300},
|
||||
A: net.ParseIP("1.2.3.4"),
|
||||
})
|
||||
|
||||
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 wildcardAnswer.Copy(), nil
|
||||
})
|
||||
|
||||
ctx := context.Background()
|
||||
results, err := tr.Traverse(ctx, "sub.example.com")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if len(results) == 0 {
|
||||
t.Fatal("expected results for wildcard response")
|
||||
}
|
||||
if results[0].Response.Type != RespAnswer {
|
||||
t.Errorf("Type = %s, want answer", results[0].Response.Type)
|
||||
}
|
||||
}
|
||||
|
||||
// TestLongCNAMEChainDepthLimit verifies that a very long CNAME chain is
|
||||
// terminated by the MaxDepth limit without infinite recursion or a panic.
|
||||
func TestLongCNAMEChainDepthLimit(t *testing.T) {
|
||||
// Every query returns a CNAME to the next label. The MaxDepth setting
|
||||
// must stop the chain.
|
||||
counter := 0
|
||||
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) {
|
||||
counter++
|
||||
q := msg.Question[0]
|
||||
resp := new(dns.Msg)
|
||||
resp.SetReply(msg)
|
||||
next := "next" + q.Name
|
||||
resp.Answer = append(resp.Answer, &dns.CNAME{
|
||||
Hdr: dns.RR_Header{Name: q.Name, Rrtype: dnsTypeCNAME, Class: dns.ClassINET},
|
||||
Target: next,
|
||||
})
|
||||
return resp, nil
|
||||
})
|
||||
|
||||
ctx := context.Background()
|
||||
results, err := tr.Traverse(ctx, "start.example.com")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if len(results) == 0 {
|
||||
t.Fatal("expected results")
|
||||
}
|
||||
// Traversal must have stopped — counter should not be unbounded.
|
||||
if counter > 50 {
|
||||
t.Errorf("too many exchange calls (%d): chain depth limit not enforced", counter)
|
||||
}
|
||||
}
|
||||
|
||||
// TestPartialBranchFailureReturnsResults verifies the graceful degradation
|
||||
// requirement: when some NS branches fail completely, the partial results from
|
||||
// successful branches are still returned.
|
||||
func TestPartialBranchFailureReturnsResults(t *testing.T) {
|
||||
// Three nameservers: first two error, third succeeds.
|
||||
referralMsg := new(dns.Msg)
|
||||
referralMsg.Rcode = dns.RcodeSuccess
|
||||
referralMsg.Authoritative = false
|
||||
referralMsg.Ns = append(referralMsg.Ns,
|
||||
&dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns1.example.com."},
|
||||
&dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns2.example.com."},
|
||||
&dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns3.example.com."},
|
||||
)
|
||||
referralMsg.Extra = append(referralMsg.Extra,
|
||||
&dns.A{Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.1")},
|
||||
&dns.A{Hdr: dns.RR_Header{Name: "ns2.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.2")},
|
||||
&dns.A{Hdr: dns.RR_Header{Name: "ns3.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.3")},
|
||||
)
|
||||
|
||||
answerMsg := new(dns.Msg)
|
||||
answerMsg.SetReply(new(dns.Msg))
|
||||
answerMsg.Answer = append(answerMsg.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) {
|
||||
switch server {
|
||||
case "198.41.0.4":
|
||||
return referralMsg.Copy(), nil
|
||||
case "10.0.0.1", "10.0.0.2":
|
||||
return nil, errors.New("server unreachable")
|
||||
case "10.0.0.3":
|
||||
return answerMsg.Copy(), nil
|
||||
}
|
||||
return nil, errors.New("unexpected server")
|
||||
})
|
||||
|
||||
ctx := context.Background()
|
||||
results, err := tr.Traverse(ctx, "example.com")
|
||||
if err != nil {
|
||||
t.Fatalf("traversal must not return a top-level error: %v", err)
|
||||
}
|
||||
|
||||
foundAnswer := false
|
||||
for _, r := range results {
|
||||
if r.Response != nil && r.Response.Type == RespAnswer {
|
||||
foundAnswer = true
|
||||
}
|
||||
}
|
||||
if !foundAnswer {
|
||||
t.Error("expected an answer from the third (reachable) nameserver despite others failing")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,187 @@
|
||||
package traverse
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"sort"
|
||||
|
||||
miekgdns "github.com/miekg/dns"
|
||||
)
|
||||
|
||||
// ServerResponse wraps a DecodedQuery for one (server, IP) query, mirroring
|
||||
// response.rb: it owns a child InfoCache seeded with the in-bailiwick records,
|
||||
// upgrades referral → referral_lame, and computes the start servers for
|
||||
// referral/restart children. The noglue/loop variants (response_noglue.rb,
|
||||
// response_loop.rb) are synthetic — no query was sent, DQ is nil.
|
||||
type ServerResponse struct {
|
||||
DQ *DecodedQuery
|
||||
Status Status
|
||||
|
||||
Qname string
|
||||
Qclass uint16
|
||||
Qtype uint16
|
||||
IP string
|
||||
Server string
|
||||
Bailiwick string
|
||||
// ParentIP is the address of the referring server; it is part of the
|
||||
// stats key for referral_lame so lame referrals from different parents
|
||||
// stay separate.
|
||||
ParentIP string
|
||||
|
||||
Cache *InfoCache
|
||||
Starters []StartServer
|
||||
StartersBailiwick string
|
||||
}
|
||||
|
||||
// NewServerResponse evaluates a decoded query in the context of parentCache.
|
||||
// server is the NS hostname that was queried (dq.IP is its address).
|
||||
func NewServerResponse(dq *DecodedQuery, server, parentIP string, parentCache *InfoCache) (*ServerResponse, error) {
|
||||
r := &ServerResponse{
|
||||
DQ: dq,
|
||||
Status: dq.Status,
|
||||
Qname: dq.Qname,
|
||||
Qclass: dq.Qclass,
|
||||
Qtype: dq.Qtype,
|
||||
IP: dq.IP,
|
||||
Server: canonicalName(server),
|
||||
Bailiwick: dq.Bailiwick,
|
||||
ParentIP: parentIP,
|
||||
Cache: NewInfoCache(parentCache),
|
||||
}
|
||||
if err := r.evaluate(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return r, nil
|
||||
}
|
||||
|
||||
// NewNoGlueResponse records a dead end: ip referred us to server inside
|
||||
// bailiwick without glue and there is no way to resolve it.
|
||||
func NewNoGlueResponse(qname string, qclass, qtype uint16, ip, server, bailiwick string) *ServerResponse {
|
||||
return &ServerResponse{
|
||||
Status: StatusNoGlue,
|
||||
Qname: canonicalName(qname),
|
||||
Qclass: qclass,
|
||||
Qtype: qtype,
|
||||
IP: ip,
|
||||
Server: canonicalName(server),
|
||||
Bailiwick: canonicalName(bailiwick),
|
||||
}
|
||||
}
|
||||
|
||||
// NewLoopResponse records a dead end: resolving server from ip would repeat
|
||||
// an ancestor referral.
|
||||
func NewLoopResponse(qname string, qclass, qtype uint16, ip, server, bailiwick string) *ServerResponse {
|
||||
return &ServerResponse{
|
||||
Status: StatusLoop,
|
||||
Qname: canonicalName(qname),
|
||||
Qclass: qclass,
|
||||
Qtype: qtype,
|
||||
IP: ip,
|
||||
Server: canonicalName(server),
|
||||
Bailiwick: canonicalName(bailiwick),
|
||||
}
|
||||
}
|
||||
|
||||
// evaluate mirrors response.rb#evaluate: cache the in-bailiwick records, then
|
||||
// for referral/restart work out the start servers from THIS branch's cache;
|
||||
// a referral whose new zone is not strictly deeper than the bailiwick is lame.
|
||||
func (r *ServerResponse) evaluate() error {
|
||||
if r.Status != StatusException {
|
||||
r.Cache.Add(r.DQ.CacheableGood)
|
||||
}
|
||||
switch r.DQ.Status {
|
||||
case StatusRestart:
|
||||
starters, bw, err := r.Cache.GetStartServers(r.DQ.Endname)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
r.Starters, r.StartersBailiwick = starters, bw
|
||||
case StatusReferral:
|
||||
starters, bw, err := r.Cache.GetStartServers(r.DQ.Endname)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
r.Starters, r.StartersBailiwick = starters, bw
|
||||
if isLameReferral(r.DQ.Bailiwick, bw) {
|
||||
r.Status = StatusReferralLame
|
||||
}
|
||||
starterNames := make([]string, len(starters))
|
||||
for i, s := range starters {
|
||||
starterNames[i] = s.Name
|
||||
}
|
||||
if !equalSorted(starterNames, r.DQ.AuthorityNames) {
|
||||
r.DQ.WarningsAdd("Referred authority names do not match query cache expectations")
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func equalSorted(a, b []string) bool {
|
||||
if len(a) != len(b) {
|
||||
return false
|
||||
}
|
||||
as := append([]string(nil), a...)
|
||||
bs := append([]string(nil), b...)
|
||||
sort.Strings(as)
|
||||
sort.Strings(bs)
|
||||
for i := range as {
|
||||
if as[i] != bs[i] {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// StatsKey is the leaf aggregation key (response.rb update_stats_key and the
|
||||
// noglue/loop variants): identical keys merge by summing probability.
|
||||
func (r *ServerResponse) StatsKey() string {
|
||||
qclass := ClassToString(r.Qclass)
|
||||
qtype := TypeToString(r.Qtype)
|
||||
switch r.Status {
|
||||
case StatusNoGlue, StatusLoop:
|
||||
return fmt.Sprintf("key:%s:%s:%s:%s:%s:%s:%s",
|
||||
r.Status, r.IP, r.Qname, qclass, qtype, r.Server, r.Bailiwick)
|
||||
default:
|
||||
key := fmt.Sprintf("key:%s:%s:%s:%s:%s:%s",
|
||||
r.Status, r.IP, r.Server, r.Qname, qclass, qtype)
|
||||
if r.Status == StatusException && r.DQ != nil {
|
||||
key += ":" + r.DQ.ExceptionMessage
|
||||
} else if r.Status == StatusReferralLame {
|
||||
key += ":" + r.ParentIP
|
||||
}
|
||||
return key
|
||||
}
|
||||
}
|
||||
|
||||
// String renders a short description for progress display (the Ruby
|
||||
// response to_s variants: "No glue for X" / "Loop encountered resolving X").
|
||||
func (r *ServerResponse) String() string {
|
||||
switch r.Status {
|
||||
case StatusNoGlue:
|
||||
return fmt.Sprintf("No glue for %s", r.Server)
|
||||
case StatusLoop:
|
||||
return fmt.Sprintf("Loop encountered resolving %s", r.Server)
|
||||
case StatusException:
|
||||
if r.DQ != nil {
|
||||
return r.DQ.ExceptionMessage
|
||||
}
|
||||
case StatusError:
|
||||
if r.DQ != nil {
|
||||
return r.DQ.ErrorMessage
|
||||
}
|
||||
}
|
||||
return string(r.Status)
|
||||
}
|
||||
|
||||
func ClassToString(qclass uint16) string {
|
||||
if s, ok := miekgdns.ClassToString[qclass]; ok {
|
||||
return s
|
||||
}
|
||||
return fmt.Sprintf("CLASS%d", qclass)
|
||||
}
|
||||
|
||||
func TypeToString(qtype uint16) string {
|
||||
if s, ok := miekgdns.TypeToString[qtype]; ok {
|
||||
return s
|
||||
}
|
||||
return fmt.Sprintf("TYPE%d", qtype)
|
||||
}
|
||||
@@ -0,0 +1,232 @@
|
||||
package traverse
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/miekg/dns"
|
||||
)
|
||||
|
||||
// rootedCache returns a cache seeded with root hints, as the traverser will
|
||||
// always provide.
|
||||
func rootedCache() *InfoCache {
|
||||
c := NewInfoCache(nil)
|
||||
c.AddHints("", []StartServer{{Name: "a.root-servers.net", IPs: []string{"198.41.0.4"}}})
|
||||
return c
|
||||
}
|
||||
|
||||
func TestServerResponseReferralNotLame(t *testing.T) {
|
||||
// Root server refers com query to the gtld servers: "" → "com" is deeper.
|
||||
msg := newMsg("www.example.com", dns.TypeA, dns.RcodeSuccess)
|
||||
msg.Ns = append(msg.Ns, nsRR("com", "a.gtld-servers.net"))
|
||||
msg.Extra = append(msg.Extra, aRR("a.gtld-servers.net", "192.5.6.30"))
|
||||
dq := decode(msg, "www.example.com", dns.TypeA, "")
|
||||
|
||||
r, err := NewServerResponse(dq, "a.root-servers.net", "", rootedCache())
|
||||
if err != nil {
|
||||
t.Fatalf("NewServerResponse: %v", err)
|
||||
}
|
||||
if r.Status != StatusReferral {
|
||||
t.Fatalf("status = %s, want referral", r.Status)
|
||||
}
|
||||
if r.StartersBailiwick != "com" {
|
||||
t.Errorf("starters bailiwick = %q, want com", r.StartersBailiwick)
|
||||
}
|
||||
if len(r.Starters) != 1 || r.Starters[0].Name != "a.gtld-servers.net" {
|
||||
t.Errorf("starters = %v", r.Starters)
|
||||
}
|
||||
if len(r.Starters[0].IPs) != 1 || r.Starters[0].IPs[0] != "192.5.6.30" {
|
||||
t.Errorf("starter IPs = %v", r.Starters[0].IPs)
|
||||
}
|
||||
if len(dq.Warnings) != 0 {
|
||||
t.Errorf("unexpected warnings: %v", dq.Warnings)
|
||||
}
|
||||
}
|
||||
|
||||
func TestServerResponseGluelessStarterHasNilIPs(t *testing.T) {
|
||||
msg := newMsg("www.example.com", dns.TypeA, dns.RcodeSuccess)
|
||||
msg.Ns = append(msg.Ns, nsRR("example.com", "ns1.example.com"))
|
||||
dq := decode(msg, "www.example.com", dns.TypeA, "com")
|
||||
|
||||
r, err := NewServerResponse(dq, "a.gtld-servers.net", "", rootedCache())
|
||||
if err != nil {
|
||||
t.Fatalf("NewServerResponse: %v", err)
|
||||
}
|
||||
if r.Status != StatusReferral {
|
||||
t.Fatalf("status = %s, want referral", r.Status)
|
||||
}
|
||||
if r.Starters[0].IPs != nil {
|
||||
t.Errorf("glueless starter should have nil IPs, got %v", r.Starters[0].IPs)
|
||||
}
|
||||
}
|
||||
|
||||
func TestServerResponseLameReferral(t *testing.T) {
|
||||
// A com server "refers" us to an example.org zone: the NS records are
|
||||
// out-of-bailiwick so they are discarded, the cache walk falls back to
|
||||
// the root NS, and "" is not strictly deeper than "com" → lame.
|
||||
msg := newMsg("www.example.com", dns.TypeA, dns.RcodeSuccess)
|
||||
msg.Ns = append(msg.Ns, nsRR("example.org", "ns1.example.org"))
|
||||
dq := decode(msg, "www.example.com", dns.TypeA, "com")
|
||||
if dq.Status != StatusReferral {
|
||||
t.Fatalf("decoded status = %s, want referral", dq.Status)
|
||||
}
|
||||
|
||||
r, err := NewServerResponse(dq, "a.gtld-servers.net", "192.5.6.30", rootedCache())
|
||||
if err != nil {
|
||||
t.Fatalf("NewServerResponse: %v", err)
|
||||
}
|
||||
if r.Status != StatusReferralLame {
|
||||
t.Fatalf("status = %s, want referral_lame", r.Status)
|
||||
}
|
||||
if r.StartersBailiwick != "" {
|
||||
t.Errorf("starters bailiwick = %q, want \"\" (root fallback)", r.StartersBailiwick)
|
||||
}
|
||||
found := false
|
||||
for _, w := range dq.Warnings {
|
||||
if w == "Referred authority names do not match query cache expectations" {
|
||||
found = true
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
t.Errorf("expected mismatch warning, got %v", dq.Warnings)
|
||||
}
|
||||
want := "key:referral_lame:192.0.2.1:a.gtld-servers.net:www.example.com:IN:A:192.5.6.30"
|
||||
if got := r.StatsKey(); got != want {
|
||||
t.Errorf("stats key = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestServerResponseEqualZoneReferralIsLame(t *testing.T) {
|
||||
// Referral back into the SAME zone (com → com) is lame: not strictly deeper.
|
||||
msg := newMsg("www.example.com", dns.TypeA, dns.RcodeSuccess)
|
||||
msg.Ns = append(msg.Ns, nsRR("com", "b.gtld-servers.net"))
|
||||
dq := decode(msg, "www.example.com", dns.TypeA, "com")
|
||||
|
||||
r, err := NewServerResponse(dq, "a.gtld-servers.net", "192.5.6.30", rootedCache())
|
||||
if err != nil {
|
||||
t.Fatalf("NewServerResponse: %v", err)
|
||||
}
|
||||
if r.Status != StatusReferralLame {
|
||||
t.Fatalf("status = %s, want referral_lame", r.Status)
|
||||
}
|
||||
}
|
||||
|
||||
func TestServerResponseRestartStarters(t *testing.T) {
|
||||
// A CNAME out of the bailiwick restarts; starters come from the deepest
|
||||
// cached zone for the new target (root here).
|
||||
msg := newMsg("www.example.com", dns.TypeA, dns.RcodeSuccess)
|
||||
msg.Answer = append(msg.Answer, cnameRR("www.example.com", "cdn.example.org"))
|
||||
dq := decode(msg, "www.example.com", dns.TypeA, "example.com")
|
||||
if dq.Status != StatusRestart {
|
||||
t.Fatalf("decoded status = %s, want restart", dq.Status)
|
||||
}
|
||||
|
||||
parent := rootedCache()
|
||||
parent.Add([]dns.RR{nsRR("example.org", "ns1.example.org"), aRR("ns1.example.org", "9.9.9.9")})
|
||||
r, err := NewServerResponse(dq, "ns1.example.com", "", parent)
|
||||
if err != nil {
|
||||
t.Fatalf("NewServerResponse: %v", err)
|
||||
}
|
||||
if r.Status != StatusRestart {
|
||||
t.Fatalf("status = %s, want restart", r.Status)
|
||||
}
|
||||
if r.StartersBailiwick != "example.org" {
|
||||
t.Errorf("starters bailiwick = %q, want example.org", r.StartersBailiwick)
|
||||
}
|
||||
if len(r.Starters) != 1 || r.Starters[0].Name != "ns1.example.org" {
|
||||
t.Errorf("starters = %v", r.Starters)
|
||||
}
|
||||
}
|
||||
|
||||
func TestServerResponseCachesGoodRecordsInChildCache(t *testing.T) {
|
||||
parent := rootedCache()
|
||||
msg := newMsg("www.example.com", dns.TypeA, dns.RcodeSuccess)
|
||||
msg.Ns = append(msg.Ns, nsRR("example.com", "ns1.example.com"))
|
||||
msg.Extra = append(msg.Extra,
|
||||
aRR("ns1.example.com", "1.2.3.4"),
|
||||
aRR("ns1.example.org", "5.6.7.8"), // out of bailiwick — discarded
|
||||
)
|
||||
dq := decode(msg, "www.example.com", dns.TypeA, "com")
|
||||
|
||||
r, err := NewServerResponse(dq, "a.gtld-servers.net", "", parent)
|
||||
if err != nil {
|
||||
t.Fatalf("NewServerResponse: %v", err)
|
||||
}
|
||||
if got := r.Cache.Get("ns1.example.com", dns.ClassINET, dns.TypeA); len(got) != 1 {
|
||||
t.Errorf("in-bailiwick glue should be cached, got %v", got)
|
||||
}
|
||||
if got := r.Cache.Get("ns1.example.org", dns.ClassINET, dns.TypeA); got != nil {
|
||||
t.Errorf("out-of-bailiwick record must be discarded, got %v", got)
|
||||
}
|
||||
// The parent cache stays clean — records live in the response's child.
|
||||
if got := parent.Get("example.com", dns.ClassINET, dns.TypeNS); got != nil {
|
||||
t.Errorf("parent cache polluted: %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestServerResponseExceptionDoesNotCache(t *testing.T) {
|
||||
dq := NewDecodedQuery(nil, errTimeout{}, "www.example.com", dns.ClassINET, dns.TypeA, "192.0.2.1", "com")
|
||||
r, err := NewServerResponse(dq, "a.gtld-servers.net", "", rootedCache())
|
||||
if err != nil {
|
||||
t.Fatalf("NewServerResponse: %v", err)
|
||||
}
|
||||
if r.Status != StatusException {
|
||||
t.Fatalf("status = %s, want exception", r.Status)
|
||||
}
|
||||
want := "key:exception:192.0.2.1:a.gtld-servers.net:www.example.com:IN:A:query timed out"
|
||||
if got := r.StatsKey(); got != want {
|
||||
t.Errorf("stats key = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
type errTimeout struct{}
|
||||
|
||||
func (errTimeout) Error() string { return "query timed out" }
|
||||
|
||||
func TestServerResponseAnsweredStatsKey(t *testing.T) {
|
||||
msg := newMsg("www.example.com", dns.TypeA, dns.RcodeSuccess)
|
||||
msg.Answer = append(msg.Answer, aRR("www.example.com", "93.184.216.34"))
|
||||
dq := decode(msg, "www.example.com", dns.TypeA, "example.com")
|
||||
|
||||
r, err := NewServerResponse(dq, "NS1.Example.Com", "", rootedCache())
|
||||
if err != nil {
|
||||
t.Fatalf("NewServerResponse: %v", err)
|
||||
}
|
||||
want := "key:answered:192.0.2.1:ns1.example.com:www.example.com:IN:A"
|
||||
if got := r.StatsKey(); got != want {
|
||||
t.Errorf("stats key = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNoGlueResponse(t *testing.T) {
|
||||
r := NewNoGlueResponse("www.example.com", dns.ClassINET, dns.TypeA, "192.5.6.30", "ns1.example.com", "example.com")
|
||||
if r.Status != StatusNoGlue {
|
||||
t.Fatalf("status = %s, want noglue", r.Status)
|
||||
}
|
||||
// NoGlue/Loop use their own field order: ip, qname, qclass, qtype, server, bailiwick.
|
||||
want := "key:noglue:192.5.6.30:www.example.com:IN:A:ns1.example.com:example.com"
|
||||
if got := r.StatsKey(); got != want {
|
||||
t.Errorf("stats key = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoopResponse(t *testing.T) {
|
||||
r := NewLoopResponse("www.example.com", dns.ClassINET, dns.TypeA, "192.5.6.30", "ns1.example.com", "example.com")
|
||||
if r.Status != StatusLoop {
|
||||
t.Fatalf("status = %s, want loop", r.Status)
|
||||
}
|
||||
want := "key:loop:192.5.6.30:www.example.com:IN:A:ns1.example.com:example.com"
|
||||
if got := r.StatsKey(); got != want {
|
||||
t.Errorf("stats key = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestServerResponseReferralNoRootHintsErrors(t *testing.T) {
|
||||
// A lame referral with a completely empty cache chain cannot compute
|
||||
// starters; the constructor surfaces the "no root hints" error.
|
||||
msg := newMsg("www.example.com", dns.TypeA, dns.RcodeSuccess)
|
||||
msg.Ns = append(msg.Ns, nsRR("example.org", "ns1.example.org"))
|
||||
dq := decode(msg, "www.example.com", dns.TypeA, "com")
|
||||
if _, err := NewServerResponse(dq, "a.gtld-servers.net", "", NewInfoCache(nil)); err == nil {
|
||||
t.Fatal("expected error when no NS reachable in cache chain")
|
||||
}
|
||||
}
|
||||
@@ -1,58 +0,0 @@
|
||||
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
|
||||
}
|
||||
@@ -1,143 +0,0 @@
|
||||
package traverse
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"gitea.hansenits.com.au/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,74 @@
|
||||
package traverse
|
||||
|
||||
import (
|
||||
"sort"
|
||||
"strings"
|
||||
|
||||
miekgdns "github.com/miekg/dns"
|
||||
)
|
||||
|
||||
// AnswerStat is one distinct answered RRset with its accumulated probability.
|
||||
type AnswerStat struct {
|
||||
// Key groups identical RRset content: sorted rdata strings joined with
|
||||
// "@@@" (summary_stats.rb get_answer_stats).
|
||||
Key string
|
||||
Prob float64
|
||||
RRs []miekgdns.RR
|
||||
}
|
||||
|
||||
// SummaryStats groups the aggregated leaves by status; answered leaves are
|
||||
// additionally grouped by RRset content, one entry per distinct RRset
|
||||
// (summary_stats.rb). The ByStatus probabilities sum to 1.0 and the answer
|
||||
// probabilities sum to ByStatus[StatusAnswered].
|
||||
type SummaryStats struct {
|
||||
ByStatus map[Status]float64
|
||||
Answers []AnswerStat
|
||||
}
|
||||
|
||||
// SummaryStats computes (and memoises) the summary grouping of this node's
|
||||
// aggregated leaf statistics (referral.rb summary_stats). It returns nil
|
||||
// until the node's statistics have been calculated.
|
||||
func (r *Referral) SummaryStats() *SummaryStats {
|
||||
if r == nil || !r.calculated || len(r.Stats) == 0 {
|
||||
return nil
|
||||
}
|
||||
if r.summaryStats != nil {
|
||||
return r.summaryStats
|
||||
}
|
||||
|
||||
stats := &SummaryStats{ByStatus: make(map[Status]float64)}
|
||||
answers := make(map[string]*AnswerStat)
|
||||
for _, leaf := range r.StatsList() {
|
||||
status := leaf.Response.Status
|
||||
stats.ByStatus[status] += leaf.Prob
|
||||
if status != StatusAnswered {
|
||||
continue
|
||||
}
|
||||
rdatas := make([]string, 0, len(leaf.Response.DQ.Answers))
|
||||
for _, rr := range leaf.Response.DQ.Answers {
|
||||
rdatas = append(rdatas, rrData(rr))
|
||||
}
|
||||
sort.Strings(rdatas)
|
||||
key := strings.Join(rdatas, "@@@")
|
||||
if e, ok := answers[key]; ok {
|
||||
e.Prob += leaf.Prob
|
||||
} else {
|
||||
answers[key] = &AnswerStat{Key: key, Prob: leaf.Prob, RRs: leaf.Response.DQ.Answers}
|
||||
}
|
||||
}
|
||||
|
||||
for _, e := range answers {
|
||||
stats.Answers = append(stats.Answers, *e)
|
||||
}
|
||||
sort.Slice(stats.Answers, func(i, j int) bool { return stats.Answers[i].Key < stats.Answers[j].Key })
|
||||
r.summaryStats = stats
|
||||
return stats
|
||||
}
|
||||
|
||||
// rrData extracts the rdata portion of a record (dnsruby rdata_to_string):
|
||||
// everything after owner/TTL/class/type in presentation format.
|
||||
func rrData(rr miekgdns.RR) string {
|
||||
s := rr.String()
|
||||
h := rr.Header().String()
|
||||
return strings.TrimPrefix(s, h)
|
||||
}
|
||||
@@ -0,0 +1,302 @@
|
||||
package traverse
|
||||
|
||||
import (
|
||||
"math"
|
||||
"net"
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/miekg/dns"
|
||||
)
|
||||
|
||||
// mockCaptureTopology reproduces the delegation behind
|
||||
// docs/captures/dnstraverse-ruby-www.example.com-A.txt: one root, thirteen
|
||||
// com gTLD servers, example.com served by two NS with three IPs each, every
|
||||
// endpoint answering the same two A records (two endpoints return them in
|
||||
// the opposite order, as in the capture).
|
||||
func mockCaptureTopology() *mockExchange {
|
||||
m := newMockExchange()
|
||||
|
||||
gtlds := []string{"a", "b", "c", "d", "e", "f", "g", "h", "i", "j", "k", "l", "m"}
|
||||
var comNS, comGlue []dns.RR
|
||||
gtldIPs := make([]string, len(gtlds))
|
||||
for i, l := range gtlds {
|
||||
ip := "192.0.2." + strconv.Itoa(i+1)
|
||||
gtldIPs[i] = ip
|
||||
comNS = append(comNS, nsRR("com", l+".gtld-servers.net"))
|
||||
comGlue = append(comGlue, aRR(l+".gtld-servers.net", ip))
|
||||
}
|
||||
m.on("202.12.27.33", "www.example.com", dns.TypeA, referralMsg(comNS, comGlue...))
|
||||
|
||||
heraIPs := []string{"108.162.192.162", "172.64.32.162", "173.245.58.162"}
|
||||
elliottIPs := []string{"108.162.195.228", "162.159.44.228", "172.64.35.228"}
|
||||
exampleReferral := referralMsg(
|
||||
[]dns.RR{
|
||||
nsRR("example.com", "hera.ns.cloudflare.com"),
|
||||
nsRR("example.com", "elliott.ns.cloudflare.com"),
|
||||
},
|
||||
aRR("hera.ns.cloudflare.com", heraIPs[0]),
|
||||
aRR("hera.ns.cloudflare.com", heraIPs[1]),
|
||||
aRR("hera.ns.cloudflare.com", heraIPs[2]),
|
||||
aRR("elliott.ns.cloudflare.com", elliottIPs[0]),
|
||||
aRR("elliott.ns.cloudflare.com", elliottIPs[1]),
|
||||
aRR("elliott.ns.cloudflare.com", elliottIPs[2]),
|
||||
)
|
||||
for _, ip := range gtldIPs {
|
||||
m.on(ip, "www.example.com", dns.TypeA, exampleReferral)
|
||||
}
|
||||
|
||||
forward := answerMsg(
|
||||
aRR("www.example.com", "104.20.23.154"),
|
||||
aRR("www.example.com", "172.66.147.243"),
|
||||
)
|
||||
reversed := answerMsg(
|
||||
aRR("www.example.com", "172.66.147.243"),
|
||||
aRR("www.example.com", "104.20.23.154"),
|
||||
)
|
||||
for _, ip := range []string{heraIPs[0], elliottIPs[0], elliottIPs[1], elliottIPs[2]} {
|
||||
m.on(ip, "www.example.com", dns.TypeA, forward)
|
||||
}
|
||||
for _, ip := range []string{heraIPs[1], heraIPs[2]} {
|
||||
m.on(ip, "www.example.com", dns.TypeA, reversed)
|
||||
}
|
||||
return m
|
||||
}
|
||||
|
||||
// TestCapturePerEndpointFractions asserts the 16.7%-per-endpoint result of
|
||||
// the www.example.com capture: 13 gTLD paths collapse (fast mode) into six
|
||||
// endpoint leaves of 1/6 each, and the summary merges the differently
|
||||
// ordered RRsets into a single 100% answered line.
|
||||
func TestCapturePerEndpointFractions(t *testing.T) {
|
||||
m := mockCaptureTopology()
|
||||
cfg := testConfig(true)
|
||||
cfg.RootAddrs = []net.IP{net.ParseIP("202.12.27.33")}
|
||||
_, root := runTraversal(t, cfg, m, "www.example.com")
|
||||
|
||||
assertSumsToOne(t, root)
|
||||
answered := leavesByStatus(root, StatusAnswered)
|
||||
if len(answered) != 6 {
|
||||
t.Fatalf("expected 6 answered leaves (one per endpoint), got %d: %v", len(answered), root.StatsList())
|
||||
}
|
||||
for _, leaf := range answered {
|
||||
if math.Abs(leaf.Prob-1.0/6) > 1e-9 {
|
||||
t.Errorf("leaf %s prob = %v, want 1/6", leaf.Key, leaf.Prob)
|
||||
}
|
||||
}
|
||||
|
||||
stats := root.SummaryStats()
|
||||
if stats == nil {
|
||||
t.Fatal("expected summary stats after calculation")
|
||||
}
|
||||
if prob := stats.ByStatus[StatusAnswered]; math.Abs(prob-1.0) > 1e-9 {
|
||||
t.Errorf("answered summary prob = %v, want 1.0", prob)
|
||||
}
|
||||
// Both RR orders share the same sorted-rdata key: one summary line.
|
||||
if len(stats.Answers) != 1 {
|
||||
t.Fatalf("expected 1 distinct answered RRset, got %d: %v", len(stats.Answers), stats.Answers)
|
||||
}
|
||||
if math.Abs(stats.Answers[0].Prob-1.0) > 1e-9 {
|
||||
t.Errorf("answer group prob = %v, want 1.0", stats.Answers[0].Prob)
|
||||
}
|
||||
if !strings.Contains(stats.Answers[0].Key, "@@@") {
|
||||
t.Errorf("answer key should join rdata with @@@, got %q", stats.Answers[0].Key)
|
||||
}
|
||||
if len(stats.Answers[0].RRs) != 2 {
|
||||
t.Errorf("answer RRs = %v", stats.Answers[0].RRs)
|
||||
}
|
||||
}
|
||||
|
||||
// TestDistinctRRsetsSeparateGroups asserts the converse of the capture test:
|
||||
// two endpoints answering DIFFERENT content produce two separate summary
|
||||
// groups, each carrying its own share of the answered probability.
|
||||
func TestDistinctRRsetsSeparateGroups(t *testing.T) {
|
||||
m := newMockExchange()
|
||||
m.on("198.41.0.4", "www.example.com", dns.TypeA, referralMsg(
|
||||
[]dns.RR{nsRR("example.com", "ns1.example.com"), nsRR("example.com", "ns2.example.com")},
|
||||
aRR("ns1.example.com", "1.1.1.1"),
|
||||
aRR("ns2.example.com", "2.2.2.2"),
|
||||
))
|
||||
// Both RRsets share their first sorted rdata so grouping must consider
|
||||
// the full content, not just the first record.
|
||||
m.on("1.1.1.1", "www.example.com", dns.TypeA, answerMsg(
|
||||
aRR("www.example.com", "1.0.0.1"),
|
||||
aRR("www.example.com", "9.9.9.9"),
|
||||
))
|
||||
m.on("2.2.2.2", "www.example.com", dns.TypeA, answerMsg(
|
||||
aRR("www.example.com", "1.0.0.1"),
|
||||
aRR("www.example.com", "8.8.8.8"),
|
||||
))
|
||||
|
||||
_, root := runTraversal(t, testConfig(false), m, "www.example.com")
|
||||
|
||||
assertSumsToOne(t, root)
|
||||
stats := root.SummaryStats()
|
||||
if stats == nil {
|
||||
t.Fatal("expected summary stats after calculation")
|
||||
}
|
||||
if prob := stats.ByStatus[StatusAnswered]; math.Abs(prob-1.0) > 1e-9 {
|
||||
t.Errorf("answered summary prob = %v, want 1.0", prob)
|
||||
}
|
||||
if len(stats.Answers) != 2 {
|
||||
t.Fatalf("expected 2 distinct answered RRsets, got %d: %v", len(stats.Answers), stats.Answers)
|
||||
}
|
||||
if stats.Answers[0].Key == stats.Answers[1].Key {
|
||||
t.Errorf("answer groups share key %q, want distinct keys", stats.Answers[0].Key)
|
||||
}
|
||||
for i, ans := range stats.Answers {
|
||||
if math.Abs(ans.Prob-0.5) > 1e-9 {
|
||||
t.Errorf("answer group %d (%q) prob = %v, want 0.5", i, ans.Key, ans.Prob)
|
||||
}
|
||||
}
|
||||
// Answers are sorted by key: "1.0.0.1@@@8.8.8.8" then "1.0.0.1@@@9.9.9.9".
|
||||
wantRdata := []string{"8.8.8.8", "9.9.9.9"}
|
||||
for i, ans := range stats.Answers {
|
||||
if !strings.Contains(ans.Key, "1.0.0.1@@@"+wantRdata[i]) {
|
||||
t.Errorf("answer group %d key = %q, want it to contain %q", i, ans.Key, "1.0.0.1@@@"+wantRdata[i])
|
||||
}
|
||||
if len(ans.RRs) != 2 {
|
||||
t.Errorf("answer group %d RRs = %v, want 2 records", i, ans.RRs)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestServfailErrorLeaf(t *testing.T) {
|
||||
m := newMockExchange()
|
||||
m.on("198.41.0.4", "www.example.com", dns.TypeA, referralMsg(
|
||||
[]dns.RR{nsRR("example.com", "ns1.example.com"), nsRR("example.com", "ns2.example.com")},
|
||||
aRR("ns1.example.com", "1.1.1.1"),
|
||||
aRR("ns2.example.com", "2.2.2.2"),
|
||||
))
|
||||
m.on("1.1.1.1", "www.example.com", dns.TypeA, answerMsg(aRR("www.example.com", "9.9.9.9")))
|
||||
m.on("2.2.2.2", "www.example.com", dns.TypeA, rcodeMsg(dns.RcodeServerFailure))
|
||||
|
||||
_, root := runTraversal(t, testConfig(false), m, "www.example.com")
|
||||
|
||||
assertSumsToOne(t, root)
|
||||
errs := leavesByStatus(root, StatusError)
|
||||
if len(errs) != 1 {
|
||||
t.Fatalf("expected 1 error leaf, got %v", root.StatsList())
|
||||
}
|
||||
if errs[0].Response.DQ.ErrorMessage != "Server failure (SERVFAIL)" {
|
||||
t.Errorf("error message = %q", errs[0].Response.DQ.ErrorMessage)
|
||||
}
|
||||
if math.Abs(errs[0].Prob-0.5) > 1e-9 {
|
||||
t.Errorf("error prob = %v, want 0.5", errs[0].Prob)
|
||||
}
|
||||
|
||||
stats := root.SummaryStats()
|
||||
if math.Abs(stats.ByStatus[StatusError]-0.5) > 1e-9 ||
|
||||
math.Abs(stats.ByStatus[StatusAnswered]-0.5) > 1e-9 {
|
||||
t.Errorf("summary by status = %v", stats.ByStatus)
|
||||
}
|
||||
total := 0.0
|
||||
for _, prob := range stats.ByStatus {
|
||||
total += prob
|
||||
}
|
||||
if math.Abs(total-1.0) > 1e-9 {
|
||||
t.Errorf("summary probabilities sum to %v, want 1.0", total)
|
||||
}
|
||||
}
|
||||
|
||||
// TestResolveSubtreeLeavesExcluded asserts that resolve-subtree leaves (the
|
||||
// A <servername> lookups for glueless NS) never reach the main aggregation:
|
||||
// they surface only through server weights.
|
||||
func TestResolveSubtreeLeavesExcluded(t *testing.T) {
|
||||
m := newMockExchange()
|
||||
m.on("198.41.0.4", "www.example.com", dns.TypeA, referralMsg(
|
||||
[]dns.RR{nsRR("com", "a.gtld-servers.net")},
|
||||
aRR("a.gtld-servers.net", "192.5.6.30"),
|
||||
))
|
||||
m.on("192.5.6.30", "www.example.com", dns.TypeA, referralMsg(
|
||||
[]dns.RR{nsRR("example.com", "ns1.example.com"), nsRR("example.com", "ns.other.net")},
|
||||
aRR("ns1.example.com", "1.1.1.1"),
|
||||
))
|
||||
m.on("1.1.1.1", "www.example.com", dns.TypeA, answerMsg(aRR("www.example.com", "9.9.9.9")))
|
||||
m.on("198.41.0.4", "ns.other.net", dns.TypeA, answerMsg(aRR("ns.other.net", "4.4.4.4")))
|
||||
m.on("4.4.4.4", "www.example.com", dns.TypeA, answerMsg(aRR("www.example.com", "9.9.9.9")))
|
||||
|
||||
_, root := runTraversal(t, testConfig(false), m, "www.example.com")
|
||||
|
||||
assertSumsToOne(t, root)
|
||||
for _, leaf := range root.StatsList() {
|
||||
if leaf.Response.Qname == "ns.other.net" {
|
||||
t.Errorf("resolve-subtree leaf leaked into main aggregation: %s", leaf.Key)
|
||||
}
|
||||
}
|
||||
if prob := root.SummaryStats().ByStatus[StatusAnswered]; math.Abs(prob-1.0) > 1e-9 {
|
||||
t.Errorf("answered summary prob = %v, want 1.0", prob)
|
||||
}
|
||||
}
|
||||
|
||||
// TestResolveFailurePseudoIPCarriesMass asserts that a failed glue resolution
|
||||
// keeps its probability: the failure becomes a "key:" pseudo-IP whose mass
|
||||
// surfaces in the main aggregation as the failing (resolve) query.
|
||||
func TestResolveFailurePseudoIPCarriesMass(t *testing.T) {
|
||||
m := newMockExchange()
|
||||
m.on("198.41.0.4", "www.example.com", dns.TypeA, referralMsg(
|
||||
[]dns.RR{nsRR("com", "a.gtld-servers.net")},
|
||||
aRR("a.gtld-servers.net", "192.5.6.30"),
|
||||
))
|
||||
m.on("192.5.6.30", "www.example.com", dns.TypeA, referralMsg(
|
||||
[]dns.RR{nsRR("example.com", "ns1.example.com"), nsRR("example.com", "ns.other.net")},
|
||||
aRR("ns1.example.com", "1.1.1.1"),
|
||||
))
|
||||
m.on("1.1.1.1", "www.example.com", dns.TypeA, answerMsg(aRR("www.example.com", "9.9.9.9")))
|
||||
// The resolve of A ns.other.net fails at the root: SERVFAIL.
|
||||
m.on("198.41.0.4", "ns.other.net", dns.TypeA, rcodeMsg(dns.RcodeServerFailure))
|
||||
|
||||
_, root := runTraversal(t, testConfig(false), m, "www.example.com")
|
||||
|
||||
assertSumsToOne(t, root)
|
||||
errs := leavesByStatus(root, StatusError)
|
||||
if len(errs) != 1 {
|
||||
t.Fatalf("expected 1 error leaf from the failed resolve, got %v", root.StatsList())
|
||||
}
|
||||
leaf := errs[0]
|
||||
if math.Abs(leaf.Prob-0.5) > 1e-9 {
|
||||
t.Errorf("failed-resolve prob = %v, want 0.5", leaf.Prob)
|
||||
}
|
||||
if leaf.Response.Qname != "ns.other.net" {
|
||||
t.Errorf("failed-resolve leaf qname = %q, want the resolve target", leaf.Response.Qname)
|
||||
}
|
||||
if !strings.HasPrefix(leaf.Key, "key:error:") {
|
||||
t.Errorf("failed-resolve key = %q", leaf.Key)
|
||||
}
|
||||
// The leaf's referral is the resolve-subtree node; its parent is the
|
||||
// glueless referral, which carries the mass as a pseudo-IP server entry.
|
||||
glueless := leaf.Referral.Parent
|
||||
if glueless.Server != "ns.other.net" {
|
||||
t.Fatalf("glueless referral server = %q", glueless.Server)
|
||||
}
|
||||
hasPseudo := false
|
||||
for ip, weight := range glueless.ServerWeights {
|
||||
if strings.HasPrefix(ip, "key:") && math.Abs(weight-1.0) <= 1e-9 {
|
||||
hasPseudo = true
|
||||
}
|
||||
}
|
||||
if !hasPseudo {
|
||||
t.Errorf("expected a key: pseudo-IP with weight 1.0, got %v", glueless.ServerWeights)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSummaryStatsNilAndMemoised(t *testing.T) {
|
||||
var nilRef *Referral
|
||||
if nilRef.SummaryStats() != nil {
|
||||
t.Error("nil referral should produce nil summary")
|
||||
}
|
||||
uncalculated := newTestReferral("ns1.example.com", []string{"1.1.1.1"})
|
||||
if uncalculated.SummaryStats() != nil {
|
||||
t.Error("uncalculated referral should produce nil summary")
|
||||
}
|
||||
|
||||
m := mockSimpleDelegation()
|
||||
_, root := runTraversal(t, testConfig(false), m, "www.example.com")
|
||||
first := root.SummaryStats()
|
||||
if first == nil {
|
||||
t.Fatal("expected summary stats")
|
||||
}
|
||||
if root.SummaryStats() != first {
|
||||
t.Error("summary stats should be memoised")
|
||||
}
|
||||
}
|
||||
@@ -1,27 +1,29 @@
|
||||
// Package traverse implements the core DNS traversal engine for ExploreDNS.
|
||||
// Package traverse implements the core DNS traversal engine for ExploreDNS,
|
||||
// a Go port of the Ruby dnstraverse engine (dns.squish.net).
|
||||
//
|
||||
// The traversal engine starts from the DNS root servers and iteratively
|
||||
// follows every referral it receives, building a complete picture of the
|
||||
// delegation path for a domain. Unlike a standard recursive resolver, which
|
||||
// stops at the first authoritative answer, the traversal engine explores every
|
||||
// branch so that delegation mismatches, lame delegations, or split authorities
|
||||
// are all visible in the output.
|
||||
// The traversal starts from a synthetic "rootroot" node (never displayed)
|
||||
// with one child per root server, and explores every branch of the
|
||||
// delegation instead of stopping at the first authoritative answer, so lame
|
||||
// delegations, missing glue and split authorities are all visible.
|
||||
//
|
||||
// # Architecture
|
||||
//
|
||||
// A Traverser maintains a stack of Referral objects. Each Referral
|
||||
// represents a pending query to a specific set of nameservers for a specific
|
||||
// name and record type. The engine pops referrals one at a time, sends the
|
||||
// query, classifies the response, and pushes any child referrals back onto the
|
||||
// stack.
|
||||
//
|
||||
// When a referral contains nameserver names but no glue records (IP addresses),
|
||||
// the engine resolves them via a secondary traversal before continuing.
|
||||
// A Traverser runs an explicit stack loop over Referral nodes. Each Referral
|
||||
// queries every IP address of one nameserver for one qname/qclass/qtype,
|
||||
// classifies each response (DecodedQuery, ServerResponse) and creates one
|
||||
// child per NS name for referral/restart statuses — including glueless
|
||||
// nameservers, which get their own resolve subtree (refid ".0." components)
|
||||
// queried from this branch's cache, never a system resolver. Post-order
|
||||
// stack markers fold the statistics upwards once all children finished:
|
||||
// every leaf outcome carries a probability, and the probabilities at the
|
||||
// root sum to 1.0.
|
||||
//
|
||||
// # Caching
|
||||
//
|
||||
// An InfoCache stores discovered glue records. In fast mode (default) a
|
||||
// single root cache is shared across all branches so that glue discovered in
|
||||
// one branch is immediately available to sibling branches. Disable fast mode
|
||||
// (TraverserConfig.Fast = false) for fully independent branch resolution.
|
||||
// Two caches cooperate: the packet-level cache in internal/dns sends each
|
||||
// (server IP, question, udpsize) at most once per run, and the hierarchical
|
||||
// per-branch InfoCache holds the in-bailiwick records each response is
|
||||
// allowed to contribute. Fast mode (default) additionally memoises completed
|
||||
// referrals so identical subtrees are reported as "completed earlier"
|
||||
// instead of being walked again.
|
||||
package traverse
|
||||
|
||||
+252
-419
@@ -4,9 +4,8 @@ import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"gitea.hansenits.com.au/hits/ExploreDNS/internal/dns"
|
||||
miekgdns "github.com/miekg/dns"
|
||||
@@ -14,495 +13,329 @@ import (
|
||||
|
||||
// TraverserConfig configures the behaviour of a Traverser.
|
||||
type TraverserConfig struct {
|
||||
// MaxDepth is the maximum referral depth before the traversal gives up.
|
||||
MaxDepth int
|
||||
// MaxDepth is the maximum referral depth (non-zero refid components)
|
||||
// before a "Maxdepth N exceeded" exception is injected.
|
||||
MaxDepth int
|
||||
// QueryType is the DNS record type to query (e.g. dns.TypeA).
|
||||
QueryType uint16
|
||||
QueryType uint16
|
||||
// RootConfig controls how root servers are discovered.
|
||||
RootConfig *dns.RootDiscoveryConfig
|
||||
RootConfig *dns.RootDiscoveryConfig
|
||||
// QueryConfig controls per-query transport parameters.
|
||||
QueryConfig *dns.QueryConfig
|
||||
// RootAddrs is an optional pre-seeded list of root server IP addresses.
|
||||
// When non-empty, root discovery via RootConfig is skipped.
|
||||
RootAddrs []net.IP
|
||||
// When non-empty, root discovery via RootConfig is skipped and each
|
||||
// address becomes one root (named by its address).
|
||||
RootAddrs []net.IP
|
||||
// Hooks provides optional callbacks for traversal events.
|
||||
Hooks *TraverserHooks
|
||||
// Fast controls cache sharing across branches. When true (default), child
|
||||
// branches inherit glue discovered by earlier branches via the shared root
|
||||
// cache, trading accuracy for speed. When false, each branch gets a
|
||||
// completely independent cache — slower but results are not contaminated by
|
||||
// sibling branch observations.
|
||||
Hooks *TraverserHooks
|
||||
// Fast enables the completed-referral memo (traverser.rb @answered):
|
||||
// a referral identical to an earlier completed one (same qname/qclass/
|
||||
// qtype/server and per-IP weights) is replaced by it instead of being
|
||||
// walked again. Non-fast mode re-walks every branch.
|
||||
Fast bool
|
||||
}
|
||||
|
||||
func DefaultTraverserConfig() *TraverserConfig {
|
||||
return &TraverserConfig{
|
||||
MaxDepth: DefaultMaxDepth,
|
||||
QueryType: dns.TypeA,
|
||||
RootConfig: nil,
|
||||
QueryConfig: nil,
|
||||
RootAddrs: nil,
|
||||
Fast: true,
|
||||
MaxDepth: DefaultMaxDepth,
|
||||
QueryType: dns.TypeA,
|
||||
Fast: true,
|
||||
}
|
||||
}
|
||||
|
||||
// TraversalResult pairs a Referral with the Response received when it was processed.
|
||||
type TraversalResult struct {
|
||||
Referral *Referral
|
||||
Response *Response
|
||||
}
|
||||
|
||||
// Traverser performs an exhaustive iterative DNS traversal starting from the
|
||||
// root servers. Create one via NewTraverser and call Traverse to start a run.
|
||||
// Traverser drives the traversal: it owns the packet-cached query client,
|
||||
// the fast-mode memo and the explicit stack loop (traverser.rb).
|
||||
type Traverser struct {
|
||||
config *TraverserConfig
|
||||
client *dns.Client
|
||||
exchange dns.ExchangeFunc
|
||||
visited map[string]bool
|
||||
depth int
|
||||
mu sync.Mutex
|
||||
// answered is the fast-mode memo of completed referrals.
|
||||
answered map[string]*Referral
|
||||
// seen maps every server name encountered to its IP addresses.
|
||||
seen map[string][]string
|
||||
// roots memoises root discovery so Roots() and Run() share one lookup.
|
||||
roots []StartServer
|
||||
}
|
||||
|
||||
func NewTraverser(cfg *TraverserConfig) *Traverser {
|
||||
if cfg == nil {
|
||||
cfg = DefaultTraverserConfig()
|
||||
}
|
||||
if cfg.MaxDepth <= 0 {
|
||||
cfg.MaxDepth = DefaultMaxDepth
|
||||
}
|
||||
if cfg.QueryType == 0 {
|
||||
cfg.QueryType = dns.TypeA
|
||||
}
|
||||
return &Traverser{
|
||||
config: cfg,
|
||||
exchange: nil,
|
||||
visited: make(map[string]bool),
|
||||
depth: 0,
|
||||
client: dns.NewClient(cfg.QueryConfig, nil),
|
||||
answered: make(map[string]*Referral),
|
||||
seen: make(map[string][]string),
|
||||
}
|
||||
}
|
||||
|
||||
// SetExchange injects a mock wire exchange into the single query path (both
|
||||
// traversal queries and root discovery); tests use this so no packets leave
|
||||
// the process.
|
||||
func (t *Traverser) SetExchange(fn dns.ExchangeFunc) {
|
||||
t.exchange = fn
|
||||
t.client = dns.NewClient(t.config.QueryConfig, fn)
|
||||
}
|
||||
|
||||
func (t *Traverser) SetHooks(hooks *TraverserHooks) {
|
||||
if t.config == nil {
|
||||
t.config = DefaultTraverserConfig()
|
||||
}
|
||||
t.config.Hooks = hooks
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
var cache *InfoCache
|
||||
if t.config.Fast {
|
||||
// Fast mode: inherit glue from the shared root cache so earlier
|
||||
// branch discoveries are visible to later branches.
|
||||
cache = rootCache
|
||||
if ref.Parent != nil {
|
||||
cache = rootCache.Child()
|
||||
}
|
||||
} else {
|
||||
// Non-fast mode: every referral gets its own independent cache so
|
||||
// no cross-branch glue is reused, ensuring each path is resolved
|
||||
// from scratch.
|
||||
cache = NewInfoCache(nil)
|
||||
}
|
||||
|
||||
if t.config.Hooks != nil {
|
||||
t.config.Hooks.emit(EventStart, TraversalResult{Referral: ref}, false)
|
||||
}
|
||||
|
||||
resp := t.processReferral(ctx, ref, cache)
|
||||
|
||||
result := TraversalResult{Referral: ref, Response: resp}
|
||||
if t.config.Hooks != nil {
|
||||
t.config.Hooks.emit(EventComplete, result, false)
|
||||
}
|
||||
|
||||
mu.Lock()
|
||||
results = append(results, result)
|
||||
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 {
|
||||
// Detect CNAME loop: target name already appears in the ancestor chain.
|
||||
if follow.Parent != nil && follow.Parent.IsNameInChain(follow.Name) {
|
||||
mu.Lock()
|
||||
results = append(results, TraversalResult{
|
||||
Referral: follow,
|
||||
Response: &Response{
|
||||
Referral: follow,
|
||||
Type: RespCNAMELoop,
|
||||
ErrorMessage: fmt.Sprintf("CNAME loop detected: %s already in traversal chain", follow.Name),
|
||||
},
|
||||
})
|
||||
mu.Unlock()
|
||||
} else if !stack.Push(follow) {
|
||||
mu.Lock()
|
||||
results = append(results, TraversalResult{
|
||||
Referral: follow,
|
||||
Response: &Response{
|
||||
Referral: follow,
|
||||
Type: RespError,
|
||||
},
|
||||
})
|
||||
mu.Unlock()
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return results, nil
|
||||
// ServersEncountered returns every server name seen during the run mapped to
|
||||
// its known IP addresses (traverser.rb servers_encountered).
|
||||
func (t *Traverser) ServersEncountered() map[string][]string {
|
||||
return t.seen
|
||||
}
|
||||
|
||||
func (t *Traverser) discoverRoots(ctx context.Context) ([]net.IP, error) {
|
||||
if len(t.config.RootAddrs) > 0 {
|
||||
return t.config.RootAddrs, nil
|
||||
// Roots performs (and memoises) root discovery, returning the start servers
|
||||
// the traversal will begin from. Callers may use it before Run to report the
|
||||
// initial root; Run reuses the memoised result.
|
||||
func (t *Traverser) Roots(ctx context.Context) ([]StartServer, error) {
|
||||
if t.roots == nil {
|
||||
roots, err := t.rootStartServers(ctx)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("root discovery: %w", err)
|
||||
}
|
||||
t.roots = roots
|
||||
}
|
||||
return t.roots, nil
|
||||
}
|
||||
|
||||
servers, err := dns.DiscoverRoots(ctx, t.config.RootConfig)
|
||||
// Run traverses the DNS for name and returns the synthetic rootroot node
|
||||
// (never displayed) whose Stats aggregate every leaf outcome; the per-leaf
|
||||
// probabilities sum to 1.0.
|
||||
func (t *Traverser) Run(ctx context.Context, name string) (*Referral, error) {
|
||||
roots, err := t.Roots(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var addrs []net.IP
|
||||
for _, srv := range servers {
|
||||
addrs = append(addrs, srv.AllIPs(false)...)
|
||||
cache := NewInfoCache(nil)
|
||||
cache.AddHints("", roots)
|
||||
|
||||
root := &Referral{
|
||||
RefID: "",
|
||||
Qname: canonicalName(toASCII(name)),
|
||||
Qclass: miekgdns.ClassINET,
|
||||
Qtype: t.config.QueryType,
|
||||
NSAType: dns.TypeA,
|
||||
Server: "",
|
||||
Bailiwick: "",
|
||||
InfoCache: cache,
|
||||
Status: RefStatusNormal,
|
||||
Responses: make(map[string]*ServerResponse),
|
||||
Children: make(map[string][]*Referral),
|
||||
ServerWeights: make(map[string]float64),
|
||||
client: t.client,
|
||||
maxdepth: t.config.MaxDepth,
|
||||
}
|
||||
return addrs, nil
|
||||
t.config.Hooks.emit(StageNew, root, "")
|
||||
|
||||
if err := t.run(ctx, root); err != nil {
|
||||
return root, err
|
||||
}
|
||||
return root, nil
|
||||
}
|
||||
|
||||
func (t *Traverser) processReferral(ctx context.Context, ref *Referral, cache *InfoCache) *Response {
|
||||
if !ref.HasAddresses() {
|
||||
t.mu.Lock()
|
||||
visitedCopy := make(map[string]bool)
|
||||
for k, v := range t.visited {
|
||||
visitedCopy[k] = v
|
||||
}
|
||||
t.mu.Unlock()
|
||||
// stack markers mirroring Ruby's :calc_resolve / :calc_answer placeholders:
|
||||
// the referral is revisited after its resolves/children finished, giving
|
||||
// post-order statistics calculation without recursion.
|
||||
type stackMarker int
|
||||
|
||||
// Resolve the nameserver's IP address. The NS hostname is stored in
|
||||
// Bailiwick; ref.Name is the domain being queried (not the NS name).
|
||||
nsToResolve := ref.Bailiwick
|
||||
if nsToResolve == "" || nsToResolve == "." {
|
||||
nsToResolve = ref.Name
|
||||
}
|
||||
nsName := strings.TrimSuffix(nsToResolve, ".")
|
||||
const (
|
||||
markerNone stackMarker = iota
|
||||
markerCalcResolve
|
||||
markerCalcAnswer
|
||||
)
|
||||
|
||||
ref.Addresses = t.resolveGlueViaSystem(ctx, nsToResolve, cache)
|
||||
if len(ref.Addresses) > 0 {
|
||||
ref.State = StateResolved
|
||||
} else {
|
||||
addrs, err := t.ResolveNS(ctx, nsToResolve, cache, visitedCopy, t.depth)
|
||||
if err != nil {
|
||||
return &Response{
|
||||
Referral: ref,
|
||||
Type: RespNSResolutionFailed,
|
||||
ErrorMessage: fmt.Sprintf("nameserver %s could not be resolved", nsName),
|
||||
}
|
||||
}
|
||||
if len(addrs) > 0 {
|
||||
ref.Addresses = addrs
|
||||
ref.State = StateResolved
|
||||
} else {
|
||||
return &Response{
|
||||
Referral: ref,
|
||||
Type: RespNSResolutionFailed,
|
||||
ErrorMessage: fmt.Sprintf("nameserver %s could not be resolved", nsName),
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
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,
|
||||
}
|
||||
type stackEntry struct {
|
||||
ref *Referral
|
||||
marker stackMarker
|
||||
}
|
||||
|
||||
func (t *Traverser) ResolveNS(ctx context.Context, nsName string, cache *InfoCache, visited map[string]bool, depth int) ([]net.IP, error) {
|
||||
if cache != nil {
|
||||
if addrs := cache.LookupGlue(nsName); len(addrs) > 0 {
|
||||
return addrs, nil
|
||||
}
|
||||
func (t *Traverser) run(ctx context.Context, root *Referral) error {
|
||||
stack := []stackEntry{{ref: root}}
|
||||
pop := func() stackEntry {
|
||||
e := stack[len(stack)-1]
|
||||
stack = stack[:len(stack)-1]
|
||||
return e
|
||||
}
|
||||
|
||||
if visited != nil {
|
||||
if visited[nsName] {
|
||||
return nil, &CircularReferralError{
|
||||
Name: nsName,
|
||||
Chain: getVisitedNames(visited),
|
||||
}
|
||||
}
|
||||
visited[nsName] = true
|
||||
}
|
||||
|
||||
if depth > DefaultMaxDepth {
|
||||
return nil, &UnresolvableNameserverError{
|
||||
Name: nsName,
|
||||
Reason: "max depth exceeded",
|
||||
}
|
||||
}
|
||||
|
||||
roots, err := t.discoverRoots(ctx)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("root discovery: %w", err)
|
||||
}
|
||||
|
||||
ref := NewReferral(nsName, dns.TypeA, ".", 0, 1.0, nil)
|
||||
ref.Addresses = roots
|
||||
ref.State = StateResolved
|
||||
|
||||
traversalCache := NewInfoCache(nil)
|
||||
if visited != nil {
|
||||
for name := range visited {
|
||||
traversalCache.StoreGlue(name, []net.IP{})
|
||||
}
|
||||
}
|
||||
|
||||
var addrs []net.IP
|
||||
var lastErr error
|
||||
|
||||
stack := NewStack(DefaultMaxDepth)
|
||||
stack.Push(ref)
|
||||
|
||||
for {
|
||||
for len(stack) > 0 {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return nil, fmt.Errorf("resolution cancelled: %w", ctx.Err())
|
||||
return fmt.Errorf("traversal cancelled: %w", ctx.Err())
|
||||
default:
|
||||
}
|
||||
|
||||
current := stack.Pop()
|
||||
if current == nil {
|
||||
break
|
||||
}
|
||||
e := pop()
|
||||
r := e.ref
|
||||
|
||||
cacheForStep := traversalCache
|
||||
if current.Parent != nil {
|
||||
cacheForStep = traversalCache.Child()
|
||||
}
|
||||
|
||||
if t.config.Hooks != nil {
|
||||
t.config.Hooks.emit(EventStart, TraversalResult{Referral: current}, true)
|
||||
}
|
||||
|
||||
resp := t.processReferral(ctx, current, cacheForStep)
|
||||
|
||||
if t.config.Hooks != nil {
|
||||
t.config.Hooks.emit(EventComplete, TraversalResult{Referral: current, Response: resp}, true)
|
||||
}
|
||||
|
||||
if resp.Type == RespAnswer && len(resp.Decoded.Answers) > 0 {
|
||||
for _, rr := range resp.Decoded.Answers {
|
||||
if a, ok := rr.(*miekgdns.A); ok {
|
||||
addrs = append(addrs, a.A)
|
||||
}
|
||||
if aaaa, ok := rr.(*miekgdns.AAAA); ok {
|
||||
addrs = append(addrs, aaaa.AAAA)
|
||||
}
|
||||
switch e.marker {
|
||||
case markerCalcResolve:
|
||||
r.resolveCalculate()
|
||||
t.config.Hooks.emit(StageResolve, r, "")
|
||||
stack = append(stack, stackEntry{ref: r}) // now needs processing
|
||||
continue
|
||||
case markerCalcAnswer:
|
||||
r.answerCalculate()
|
||||
t.config.Hooks.emit(StageAnswer, r, "")
|
||||
if t.config.Fast && r.Status == RefStatusNormal && !hasLameResponse(r) {
|
||||
t.answered[fastKey(r)] = r
|
||||
}
|
||||
if len(addrs) > 0 {
|
||||
if cache != nil {
|
||||
cache.StoreGlue(nsName, addrs)
|
||||
}
|
||||
return addrs, nil
|
||||
if !r.IsRootRoot() {
|
||||
t.recordSeen(r)
|
||||
}
|
||||
}
|
||||
|
||||
if resp.Type == RespNXDOMAIN {
|
||||
lastErr = &UnresolvableNameserverError{
|
||||
Name: nsName,
|
||||
Reason: "NXDOMAIN",
|
||||
}
|
||||
break
|
||||
}
|
||||
|
||||
if resp.Type == RespSERVFAIL || resp.Type == RespError || resp.Type == RespNSResolutionFailed {
|
||||
lastErr = fmt.Errorf("server error resolving %s: %s", nsName, resp.Type)
|
||||
continue
|
||||
}
|
||||
|
||||
if resp.Type == RespReferral {
|
||||
children := resp.ChildReferrals()
|
||||
for _, child := range children {
|
||||
// Only skip visited names when they have no addresses; if glue
|
||||
// was included in the referral response we still need to query
|
||||
// that child to get the authoritative answer.
|
||||
if visited != nil && visited[child.Name] && !child.HasAddresses() {
|
||||
continue
|
||||
// A new item. Fast mode: an identical completed referral replaces
|
||||
// this one wholesale. noglue/loop nodes are excluded because their
|
||||
// stats carry node-specific attributes and are cheap to recreate.
|
||||
if t.config.Fast && r.Parent != nil {
|
||||
if memo, ok := t.answered[fastKey(r)]; ok && !r.isNoGlue() && !r.isLoop() {
|
||||
r.Parent.replaceChild(r, memo)
|
||||
t.config.Hooks.emit(StageAnswerFast, r, memo.RefID)
|
||||
continue
|
||||
}
|
||||
}
|
||||
|
||||
t.config.Hooks.emit(StageStart, r, "")
|
||||
|
||||
if !r.Resolved() {
|
||||
// Push the resolve subtree with a calc_resolve placeholder so the
|
||||
// weights are folded in once every resolve leaf completed.
|
||||
stack = append(stack, stackEntry{ref: r, marker: markerCalcResolve})
|
||||
resolves, err := r.resolve()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
for _, c := range resolves {
|
||||
t.config.Hooks.emit(StageNew, c, "")
|
||||
}
|
||||
for i := len(resolves) - 1; i >= 0; i-- {
|
||||
stack = append(stack, stackEntry{ref: resolves[i]})
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
stack = append(stack, stackEntry{ref: r, marker: markerCalcAnswer})
|
||||
childrenSets, err := r.process(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
seenParentIP := make(map[string]bool)
|
||||
var flat []*Referral
|
||||
for _, set := range childrenSets {
|
||||
for _, c := range set {
|
||||
if len(childrenSets) > 1 && !seenParentIP[c.ParentIP] {
|
||||
t.config.Hooks.emit(StageNewReferralSet, c, "")
|
||||
seenParentIP[c.ParentIP] = true
|
||||
}
|
||||
if !stack.Push(child) {
|
||||
lastErr = &UnresolvableNameserverError{
|
||||
Name: nsName,
|
||||
Reason: "max depth exceeded during resolution",
|
||||
stage, earlier := StageNew, ""
|
||||
if t.config.Fast {
|
||||
if memo, ok := t.answered[fastKey(c)]; ok {
|
||||
stage, earlier = StageNewFast, memo.RefID
|
||||
}
|
||||
}
|
||||
t.config.Hooks.emit(stage, c, earlier)
|
||||
flat = append(flat, c)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if len(addrs) > 0 {
|
||||
return addrs, nil
|
||||
}
|
||||
|
||||
if lastErr != nil {
|
||||
return nil, lastErr
|
||||
}
|
||||
|
||||
return nil, &UnresolvableNameserverError{
|
||||
Name: nsName,
|
||||
Reason: "resolution exhausted without answer",
|
||||
}
|
||||
}
|
||||
|
||||
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)
|
||||
for i := len(flat) - 1; i >= 0; i-- {
|
||||
stack = append(stack, stackEntry{ref: flat[i]})
|
||||
}
|
||||
}
|
||||
|
||||
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 * time.Second,
|
||||
WriteTimeout: 5 * time.Second,
|
||||
}
|
||||
if deadline, ok := ctx.Deadline(); ok {
|
||||
remaining := time.Until(deadline)
|
||||
if remaining <= 0 {
|
||||
return nil
|
||||
}
|
||||
c.ReadTimeout = remaining
|
||||
c.WriteTimeout = remaining
|
||||
}
|
||||
|
||||
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
|
||||
// fastKey is the fast-mode memo key (traverser.rb): qname/qclass/qtype/
|
||||
// server plus the per-IP weights, lowercased.
|
||||
func fastKey(r *Referral) string {
|
||||
return strings.ToLower(fmt.Sprintf("%s:%s:%s:%s:%s",
|
||||
r.Qname, ClassToString(r.Qclass), TypeToString(r.Qtype), r.Server, r.TxtIPsVerbose()))
|
||||
}
|
||||
|
||||
func hasLameResponse(r *Referral) bool {
|
||||
for _, resp := range r.Responses {
|
||||
if resp.Status == StatusReferralLame {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (t *Traverser) recordSeen(r *Referral) {
|
||||
name := strings.ToLower(r.Server)
|
||||
existing := t.seen[name]
|
||||
for _, ip := range r.IPsAsArray() {
|
||||
found := false
|
||||
for _, have := range existing {
|
||||
if have == ip {
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
existing = append(existing, ip)
|
||||
}
|
||||
}
|
||||
t.seen[name] = existing
|
||||
}
|
||||
|
||||
// rootStartServers returns the root servers as start-server hints: either
|
||||
// the pre-seeded RootAddrs or the servers found via root discovery (one by
|
||||
// default, all of them with AllRoots). IPv4 only, like the reference.
|
||||
func (t *Traverser) rootStartServers(ctx context.Context) ([]StartServer, error) {
|
||||
if len(t.config.RootAddrs) > 0 {
|
||||
var out []StartServer
|
||||
for _, ip := range t.config.RootAddrs {
|
||||
if v4 := ip.To4(); v4 != nil {
|
||||
out = append(out, StartServer{Name: v4.String(), IPs: []string{v4.String()}})
|
||||
}
|
||||
}
|
||||
if len(out) == 0 {
|
||||
return nil, fmt.Errorf("no usable IPv4 root addresses")
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
rootCfg := t.config.RootConfig
|
||||
if t.exchange != nil {
|
||||
var cp dns.RootDiscoveryConfig
|
||||
if rootCfg != nil {
|
||||
cp = *rootCfg
|
||||
}
|
||||
cp.Exchange = t.exchange
|
||||
rootCfg = &cp
|
||||
}
|
||||
|
||||
servers, err := dns.DiscoverRoots(ctx, rootCfg)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var out []StartServer
|
||||
for _, srv := range servers {
|
||||
var ips []string
|
||||
for _, ip := range srv.IPv4 {
|
||||
ips = append(ips, ip.String())
|
||||
}
|
||||
if len(ips) == 0 {
|
||||
continue
|
||||
}
|
||||
out = append(out, StartServer{Name: canonicalName(srv.Name), IPs: ips})
|
||||
}
|
||||
if len(out) == 0 {
|
||||
return nil, fmt.Errorf("no root servers with IPv4 addresses")
|
||||
}
|
||||
sort.Slice(out, func(i, j int) bool { return out[i].Name < out[j].Name })
|
||||
return out, nil
|
||||
}
|
||||
|
||||
+725
-433
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user