From 39529dacbb80f5e6687a460d57f303472713fc6c Mon Sep 17 00:00:00 2001 From: Hansen IT Solutions Date: Sat, 6 Jun 2026 03:41:14 +0000 Subject: [PATCH] feat: implement traversal engine (HAN-380) (#6) Co-authored-by: Hansen IT Solutions Co-committed-by: Hansen IT Solutions --- internal/dns/query.go | 68 ++++ internal/traverse/cache.go | 124 +++++++ internal/traverse/cache_test.go | 203 ++++++++++++ internal/traverse/referral.go | 78 +++++ internal/traverse/referral_test.go | 111 +++++++ internal/traverse/response.go | 218 +++++++++++++ internal/traverse/response_test.go | 338 +++++++++++++++++++ internal/traverse/stack.go | 58 ++++ internal/traverse/stack_test.go | 143 ++++++++ internal/traverse/traverser.go | 274 ++++++++++++++++ internal/traverse/traverser_test.go | 483 ++++++++++++++++++++++++++++ 11 files changed, 2098 insertions(+) create mode 100644 internal/traverse/cache.go create mode 100644 internal/traverse/cache_test.go create mode 100644 internal/traverse/referral.go create mode 100644 internal/traverse/referral_test.go create mode 100644 internal/traverse/response.go create mode 100644 internal/traverse/response_test.go create mode 100644 internal/traverse/stack.go create mode 100644 internal/traverse/stack_test.go create mode 100644 internal/traverse/traverser.go create mode 100644 internal/traverse/traverser_test.go diff --git a/internal/dns/query.go b/internal/dns/query.go index e219f8a..ace16ee 100644 --- a/internal/dns/query.go +++ b/internal/dns/query.go @@ -125,6 +125,74 @@ func QueryWithExchange(ctx context.Context, server net.IP, name string, qtype ui return nil, fmt.Errorf("query %s %s failed after %d retries: %w", name, QNameType(qtype), cfg.Retries, lastErr) } +func IterativeQuery(ctx context.Context, server net.IP, name string, qtype uint16, cfg *QueryConfig) (*dns.Msg, error) { + if cfg == nil { + cfg = DefaultQueryConfig() + } + if cfg.UDPSize <= 0 { + cfg.UDPSize = DefaultEDNS0UDPSize() + } + return IterativeQueryWithExchange(ctx, server, name, qtype, cfg, realExchange) +} + +func IterativeQueryWithExchange(ctx context.Context, server net.IP, name string, qtype uint16, cfg *QueryConfig, exchangeFn ExchangeFunc) (*dns.Msg, error) { + if cfg == nil { + cfg = DefaultQueryConfig() + } + if cfg.UDPSize <= 0 { + cfg.UDPSize = DefaultEDNS0UDPSize() + } + + msg := buildQuery(name, qtype, cfg.UDPSize) + msg.RecursionDesired = false + serverStr := server.String() + + var lastErr error + + for attempt := 0; attempt < cfg.Retries; attempt++ { + if attempt > 0 { + select { + case <-ctx.Done(): + return nil, fmt.Errorf("query retries cancelled: %w", ctx.Err()) + case <-time.After(100 * time.Millisecond): + } + } + + if cfg.UseTCP { + resp, err := exchangeFn(ctx, serverStr, msg, true) + if err != nil { + lastErr = err + continue + } + return resp, nil + } + + resp, err := exchangeFn(ctx, serverStr, msg, false) + if err != nil { + lastErr = err + continue + } + + if resp == nil { + lastErr = fmt.Errorf("nil response") + continue + } + + if resp.Truncated { + resp, err = exchangeFn(ctx, serverStr, msg, true) + if err != nil { + lastErr = err + continue + } + return resp, nil + } + + return resp, nil + } + + return nil, fmt.Errorf("iterative query %s %s failed after %d retries: %w", name, QNameType(qtype), cfg.Retries, lastErr) +} + func buildQuery(name string, qtype uint16, udpSize int) *dns.Msg { m := new(dns.Msg) m.SetQuestion(dns.Fqdn(name), qtype) diff --git a/internal/traverse/cache.go b/internal/traverse/cache.go new file mode 100644 index 0000000..e4422ae --- /dev/null +++ b/internal/traverse/cache.go @@ -0,0 +1,124 @@ +package traverse + +import ( + "net" + "strings" + "sync" + + miekgdns "github.com/miekg/dns" +) + +type InfoCache struct { + parent *InfoCache + mu sync.RWMutex + ns map[string][]string + glue map[string][]net.IP +} + +func NewInfoCache(parent *InfoCache) *InfoCache { + return &InfoCache{ + parent: parent, + ns: make(map[string][]string), + glue: make(map[string][]net.IP), + } +} + +func (c *InfoCache) StoreNS(zone string, nameservers []string) { + if len(nameservers) == 0 { + return + } + zone = normalize(zone) + c.mu.Lock() + seen := make(map[string]bool) + for _, ns := range nameservers { + ns = normalize(ns) + if !seen[ns] { + seen[ns] = true + c.ns[zone] = append(c.ns[zone], ns) + } + } + c.mu.Unlock() +} + +func (c *InfoCache) LookupNS(zone string) []string { + zone = normalize(zone) + if names := c.localNS(zone); len(names) > 0 { + return names + } + if c.parent != nil { + return c.parent.LookupNS(zone) + } + return nil +} + +func (c *InfoCache) localNS(zone string) []string { + c.mu.RLock() + defer c.mu.RUnlock() + names, ok := c.ns[zone] + if !ok { + return nil + } + result := make([]string, len(names)) + copy(result, names) + return result +} + +func (c *InfoCache) StoreGlue(name string, addrs []net.IP) { + if len(addrs) == 0 { + return + } + name = normalize(name) + c.mu.Lock() + seen := make(map[string]bool) + for _, addr := range addrs { + key := addr.String() + if !seen[key] { + seen[key] = true + c.glue[name] = append(c.glue[name], addr) + } + } + c.mu.Unlock() +} + +func (c *InfoCache) LookupGlue(name string) []net.IP { + name = normalize(name) + if addrs := c.localGlue(name); len(addrs) > 0 { + return addrs + } + if c.parent != nil { + return c.parent.LookupGlue(name) + } + return nil +} + +func (c *InfoCache) localGlue(name string) []net.IP { + c.mu.RLock() + defer c.mu.RUnlock() + addrs, ok := c.glue[name] + if !ok { + return nil + } + result := make([]net.IP, len(addrs)) + copy(result, addrs) + return result +} + +func (c *InfoCache) Child() *InfoCache { + return NewInfoCache(c) +} + +func (c *InfoCache) NSCount() int { + c.mu.RLock() + defer c.mu.RUnlock() + return len(c.ns) +} + +func (c *InfoCache) GlueCount() int { + c.mu.RLock() + defer c.mu.RUnlock() + return len(c.glue) +} + +func normalize(name string) string { + return strings.ToLower(miekgdns.Fqdn(name)) +} diff --git a/internal/traverse/cache_test.go b/internal/traverse/cache_test.go new file mode 100644 index 0000000..cc024fa --- /dev/null +++ b/internal/traverse/cache_test.go @@ -0,0 +1,203 @@ +package traverse + +import ( + "fmt" + "net" + "strings" + "sync" + "testing" + + "github.com/miekg/dns" +) + +func TestNewInfoCache(t *testing.T) { + c := NewInfoCache(nil) + if c.parent != nil { + t.Error("root cache should have nil parent") + } + if c.NSCount() != 0 { + t.Errorf("NSCount = %d, want 0", c.NSCount()) + } + if c.GlueCount() != 0 { + t.Errorf("GlueCount = %d, want 0", c.GlueCount()) + } +} + +func TestInfoCacheStoreAndLookupNS(t *testing.T) { + c := NewInfoCache(nil) + + c.StoreNS("com.", []string{"a.gtld-servers.net.", "b.gtld-servers.net."}) + if c.NSCount() != 1 { + t.Errorf("NSCount = %d, want 1", c.NSCount()) + } + + names := c.LookupNS("com.") + if len(names) != 2 { + t.Fatalf("expected 2 nameservers, got %d", len(names)) + } + if names[0] != "a.gtld-servers.net." { + t.Errorf("nameserver[0] = %q, want %q", names[0], "a.gtld-servers.net.") + } +} + +func TestInfoCacheNSDedup(t *testing.T) { + c := NewInfoCache(nil) + c.StoreNS("com.", []string{"a.gtld-servers.net.", "a.gtld-servers.net."}) + names := c.LookupNS("com.") + if len(names) != 1 { + t.Errorf("expected 1 deduped NS, got %d", len(names)) + } +} + +func TestInfoCacheNSCaseInsensitive(t *testing.T) { + c := NewInfoCache(nil) + c.StoreNS("COM.", []string{"A.GTLD-SERVERS.NET."}) + names := c.LookupNS("com.") + if len(names) != 1 { + t.Fatalf("expected 1 NS, got %d", len(names)) + } + if names[0] != "a.gtld-servers.net." { + t.Errorf("NS = %q, want %q", names[0], "a.gtld-servers.net.") + } +} + +func TestInfoCacheNSLookupMiss(t *testing.T) { + c := NewInfoCache(nil) + names := c.LookupNS("org.") + if names != nil { + t.Errorf("expected nil for miss, got %v", names) + } +} + +func TestInfoCacheNSStoreEmpty(t *testing.T) { + c := NewInfoCache(nil) + c.StoreNS("com.", nil) + if c.NSCount() != 0 { + t.Errorf("expected 0 after empty store, got %d", c.NSCount()) + } +} + +func TestInfoCacheChainedNS(t *testing.T) { + parent := NewInfoCache(nil) + parent.StoreNS("com.", []string{"a.gtld-servers.net."}) + + child := parent.Child() + if child.parent != parent { + t.Error("child parent should be the parent cache") + } + + names := child.LookupNS("com.") + if len(names) != 1 { + t.Fatalf("expected 1 NS from parent, got %d", len(names)) + } + if child.NSCount() != 0 { + t.Errorf("child NSCount = %d, want 0", child.NSCount()) + } +} + +func TestInfoCacheChildOverridesParent(t *testing.T) { + parent := NewInfoCache(nil) + parent.StoreNS("com.", []string{"a.gtld-servers.net."}) + + child := parent.Child() + child.StoreNS("com.", []string{"b.gtld-servers.net."}) + + names := child.LookupNS("com.") + if len(names) != 1 { + t.Fatalf("expected 1 NS, got %d", len(names)) + } + if names[0] != "b.gtld-servers.net." { + t.Errorf("expected child's NS to override, got %q", names[0]) + } +} + +func TestInfoCacheStoreAndLookupGlue(t *testing.T) { + c := NewInfoCache(nil) + addrs := []net.IP{net.ParseIP("1.2.3.4"), net.ParseIP("5.6.7.8")} + c.StoreGlue("ns1.example.com.", addrs) + + result := c.LookupGlue("ns1.example.com.") + if len(result) != 2 { + t.Fatalf("expected 2 glue addresses, got %d", len(result)) + } +} + +func TestInfoCacheGlueDedup(t *testing.T) { + c := NewInfoCache(nil) + ip := net.ParseIP("1.2.3.4") + c.StoreGlue("ns1.example.com.", []net.IP{ip, ip}) + result := c.LookupGlue("ns1.example.com.") + if len(result) != 1 { + t.Errorf("expected 1 deduped glue, got %d", len(result)) + } +} + +func TestInfoCacheChainedGlue(t *testing.T) { + parent := NewInfoCache(nil) + parent.StoreGlue("ns1.example.com.", []net.IP{net.ParseIP("1.2.3.4")}) + + child := parent.Child() + result := child.LookupGlue("ns1.example.com.") + if len(result) != 1 { + t.Fatalf("expected 1 glue from parent, got %d", len(result)) + } + if child.GlueCount() != 0 { + t.Errorf("child GlueCount = %d, want 0", child.GlueCount()) + } +} + +func TestInfoCacheGlueLookupMiss(t *testing.T) { + c := NewInfoCache(nil) + result := c.LookupGlue("nonexistent.example.com.") + if result != nil { + t.Errorf("expected nil for miss, got %v", result) + } +} + +func TestInfoCacheConcurrentAccess(t *testing.T) { + c := NewInfoCache(nil) + var wg sync.WaitGroup + + for i := 0; i < 100; i++ { + wg.Add(1) + go func(i int) { + defer wg.Done() + name := strings.ToLower(dns.Fqdn(fmt.Sprintf("ns%d.example.com.", i))) + c.StoreNS("example.com.", []string{name}) + c.StoreGlue(name, []net.IP{net.ParseIP(fmt.Sprintf("1.2.3.%d", i%256))}) + _ = c.LookupNS("example.com.") + _ = c.LookupGlue(name) + }(i) + } + wg.Wait() +} + +func TestInfoCacheNilParent(t *testing.T) { + c := NewInfoCache(nil) + if c.LookupNS("com.") != nil { + t.Error("root cache should return nil for miss") + } + if c.LookupGlue("ns.example.com.") != nil { + t.Error("root cache should return nil for glue miss") + } +} + +func TestNormalize(t *testing.T) { + tests := []struct { + input string + want string + }{ + {"example.com", "example.com."}, + {"Example.COM.", "example.com."}, + {"EXAMPLE.COM", "example.com."}, + } + + for _, tt := range tests { + t.Run(tt.input, func(t *testing.T) { + got := normalize(tt.input) + if got != tt.want { + t.Errorf("normalize(%q) = %q, want %q", tt.input, got, tt.want) + } + }) + } +} diff --git a/internal/traverse/referral.go b/internal/traverse/referral.go new file mode 100644 index 0000000..b6a553f --- /dev/null +++ b/internal/traverse/referral.go @@ -0,0 +1,78 @@ +package traverse + +import ( + "net" + "strings" + + miekgdns "github.com/miekg/dns" +) + +type ResolutionState int + +const ( + StateUnresolved ResolutionState = iota + StateResolving + StateResolved +) + +func (s ResolutionState) String() string { + switch s { + case StateUnresolved: + return "unresolved" + case StateResolving: + return "resolving" + case StateResolved: + return "resolved" + default: + return "unknown" + } +} + +type Referral struct { + Name string + Qtype uint16 + Qclass uint16 + Bailiwick string + + Addresses []net.IP + State ResolutionState + + NSName string + Parent *Referral + Depth int + Prob float64 +} + +func NewReferral(name string, qtype uint16, bailiwick string, depth int, prob float64, parent *Referral) *Referral { + return &Referral{ + Name: miekgdns.Fqdn(strings.ToLower(name)), + Qtype: qtype, + Qclass: miekgdns.ClassINET, + Bailiwick: miekgdns.Fqdn(strings.ToLower(bailiwick)), + Depth: depth, + Prob: prob, + Parent: parent, + State: StateUnresolved, + } +} + +func (r *Referral) InBailiwick(name string) bool { + if r.Bailiwick == "" || r.Bailiwick == "." { + return true + } + fqdn := miekgdns.Fqdn(strings.ToLower(name)) + return miekgdns.IsSubDomain(r.Bailiwick, fqdn) +} + +func (r *Referral) HasAddresses() bool { + return len(r.Addresses) > 0 +} + +func (r *Referral) SetAddresses(addrs []net.IP) { + r.Addresses = addrs + if len(addrs) > 0 { + r.State = StateResolved + } else { + r.State = StateUnresolved + } +} diff --git a/internal/traverse/referral_test.go b/internal/traverse/referral_test.go new file mode 100644 index 0000000..572d12d --- /dev/null +++ b/internal/traverse/referral_test.go @@ -0,0 +1,111 @@ +package traverse + +import ( + "net" + "testing" + + "github.com/hits/ExploreDNS/internal/dns" +) + +const TypeA = dns.TypeA + +type State = ResolutionState + +func TestNewReferral(t *testing.T) { + ref := NewReferral("example.com", TypeA, ".", 0, 1.0, nil) + if ref.Name != "example.com." { + t.Errorf("Name = %q, want %q", ref.Name, "example.com.") + } + if ref.Qtype != TypeA { + t.Errorf("Qtype = %d, want %d", ref.Qtype, TypeA) + } + if ref.Qclass != 1 { + t.Errorf("Qclass = %d, want 1", ref.Qclass) + } + if ref.State != StateUnresolved { + t.Errorf("State = %d, want %d", ref.State, StateUnresolved) + } + if ref.Depth != 0 { + t.Errorf("Depth = %d, want 0", ref.Depth) + } + if ref.Prob != 1.0 { + t.Errorf("Prob = %f, want 1.0", ref.Prob) + } + if ref.Parent != nil { + t.Error("Parent should be nil") + } +} + +func TestReferralInBailiwick(t *testing.T) { + tests := []struct { + name string + bailiwick string + testName string + want bool + }{ + {"root bailiwick accepts all", ".", "example.com.", true}, + {"empty bailiwick accepts all", "", "example.com.", true}, + {"subdomain in bailiwick", "com.", "example.com.", true}, + {"deeper subdomain", "com.", "www.example.com.", true}, + {"not in bailiwick", "org.", "example.com.", false}, + {"same zone", "example.com.", "example.com.", true}, + {"sibling zone", "example.com.", "other.com.", false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ref := &Referral{Bailiwick: tt.bailiwick} + if got := ref.InBailiwick(tt.testName); got != tt.want { + t.Errorf("InBailiwick(%q) = %v, want %v", tt.testName, got, tt.want) + } + }) + } +} + +func TestReferralHasAddresses(t *testing.T) { + ref := &Referral{} + if ref.HasAddresses() { + t.Error("empty referral should not have addresses") + } + + ref.Addresses = []net.IP{net.ParseIP("1.2.3.4")} + if !ref.HasAddresses() { + t.Error("referral with address should have addresses") + } +} + +func TestReferralSetAddresses(t *testing.T) { + ref := &Referral{} + + ref.SetAddresses([]net.IP{net.ParseIP("1.2.3.4")}) + if ref.State != StateResolved { + t.Errorf("State = %d, want %d", ref.State, StateResolved) + } + if !ref.HasAddresses() { + t.Error("should have addresses after SetAddresses") + } + + ref.SetAddresses(nil) + if ref.State != StateUnresolved { + t.Errorf("State = %d, want %d", ref.State, StateUnresolved) + } +} + +func TestResolutionStateString(t *testing.T) { + tests := []struct { + state State + want string + }{ + {StateUnresolved, "unresolved"}, + {StateResolving, "resolving"}, + {StateResolved, "resolved"}, + } + + for _, tt := range tests { + t.Run(tt.want, func(t *testing.T) { + if got := tt.state.String(); got != tt.want { + t.Errorf("String() = %q, want %q", got, tt.want) + } + }) + } +} diff --git a/internal/traverse/response.go b/internal/traverse/response.go new file mode 100644 index 0000000..2176a0a --- /dev/null +++ b/internal/traverse/response.go @@ -0,0 +1,218 @@ +package traverse + +import ( + "net" + + "github.com/hits/ExploreDNS/internal/dns" + miekgdns "github.com/miekg/dns" +) + +type ResponseType int + +const ( + RespReferral ResponseType = iota + RespAnswer + RespCNAMEFollow + RespNODATA + RespNXDOMAIN + RespSERVFAIL + RespError +) + +func (rt ResponseType) String() string { + switch rt { + case RespReferral: + return "referral" + case RespAnswer: + return "answer" + case RespCNAMEFollow: + return "cname_follow" + case RespNODATA: + return "nodata" + case RespNXDOMAIN: + return "nxdomain" + case RespSERVFAIL: + return "servfail" + case RespError: + return "error" + default: + return "unknown" + } +} + +type Response struct { + Referral *Referral + Server net.IP + Cache *InfoCache + Decoded *dns.DecodedResponse + Type ResponseType +} + +func NewResponse(ref *Referral, server net.IP, cache *InfoCache) *Response { + return &Response{ + Referral: ref, + Server: server, + Cache: cache, + } +} + +func (r *Response) Process(msg *miekgdns.Msg) *Response { + if msg == nil { + r.Type = RespError + return r + } + + r.Decoded = dns.DecodeResponse(msg) + if r.Decoded == nil { + r.Type = RespError + return r + } + + r.Type = r.classify() + return r +} + +func (r *Response) classify() ResponseType { + switch r.Decoded.Classification { + case dns.ResponseNXDOMAIN: + return RespNXDOMAIN + case dns.ResponseSERVFAIL: + return RespSERVFAIL + case dns.ResponseAnswer: + if len(r.Decoded.CNAMEChain) > 0 && !r.hasFinalAnswer() { + return RespCNAMEFollow + } + return RespAnswer + case dns.ResponseReferral: + return RespReferral + case dns.ResponseNODATA: + return RespNODATA + default: + return RespError + } +} + +func (r *Response) hasFinalAnswer() bool { + for _, rr := range r.Decoded.Answers { + if _, ok := rr.(*miekgdns.CNAME); ok { + continue + } + return true + } + return false +} + +func (r *Response) ChildReferrals() []*Referral { + if r.Type != RespReferral { + return nil + } + if r.Referral == nil { + return nil + } + + var nameservers []string + for _, rr := range r.Decoded.Authority { + if ns, ok := rr.(*miekgdns.NS); ok { + if r.Referral.InBailiwick(ns.Ns) { + nameservers = append(nameservers, ns.Ns) + } + } + } + + if len(nameservers) == 0 { + for _, rr := range r.Decoded.Authority { + if ns, ok := rr.(*miekgdns.NS); ok { + nameservers = append(nameservers, ns.Ns) + } + } + } + + r.storeAuthority(nameservers) + + prob := r.childProb(len(nameservers)) + var children []*Referral + for _, ns := range nameservers { + child := NewReferral( + r.Referral.Name, + r.Referral.Qtype, + ns, + r.Referral.Depth+1, + prob, + r.Referral, + ) + r.resolveGlue(child) + children = append(children, child) + } + return children +} + +func (r *Response) CNAMEFollowReferral() *Referral { + if r.Type != RespCNAMEFollow || len(r.Decoded.CNAMEChain) == 0 { + return nil + } + target := r.Decoded.CNAMEChain[len(r.Decoded.CNAMEChain)-1] + follow := NewReferral( + target, + r.Referral.Qtype, + r.Referral.Bailiwick, + r.Referral.Depth+1, + r.Referral.Prob, + r.Referral, + ) + if len(r.Referral.Addresses) > 0 { + follow.Addresses = make([]net.IP, len(r.Referral.Addresses)) + copy(follow.Addresses, r.Referral.Addresses) + follow.State = StateResolved + } + return follow +} + +func (r *Response) storeAuthority(nameservers []string) { + if r.Cache == nil { + return + } + zone := r.Referral.Name + r.Cache.StoreNS(zone, nameservers) +} + +func (r *Response) resolveGlue(child *Referral) { + if r.Cache == nil { + return + } + nsName := child.Bailiwick + for _, rr := range r.Decoded.Additional { + switch v := rr.(type) { + case *miekgdns.A: + if normalize(v.Header().Name) == normalize(nsName) { + child.Addresses = append(child.Addresses, v.A) + } + case *miekgdns.AAAA: + if normalize(v.Header().Name) == normalize(nsName) { + child.Addresses = append(child.Addresses, v.AAAA) + } + } + } + if child.HasAddresses() { + child.State = StateResolved + } + r.Cache.StoreGlue(nsName, child.Addresses) +} + +func (r *Response) IsTerminal() bool { + switch r.Type { + case RespAnswer, RespNODATA, RespNXDOMAIN, RespSERVFAIL, RespError: + return true + default: + return false + } +} + +func (r *Response) childProb(n int) float64 { + if n <= 0 { + return 0 + } + if r.Referral == nil { + return 1.0 / float64(n) + } + return r.Referral.Prob / float64(n) +} diff --git a/internal/traverse/response_test.go b/internal/traverse/response_test.go new file mode 100644 index 0000000..efdcc69 --- /dev/null +++ b/internal/traverse/response_test.go @@ -0,0 +1,338 @@ +package traverse + +import ( + "net" + "testing" + + "github.com/miekg/dns" +) + +func TestResponseProcessNil(t *testing.T) { + ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil) + r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil) + r.Process(nil) + if r.Type != RespError { + t.Errorf("Type = %d, want %d", r.Type, RespError) + } +} + +func TestResponseClassifyAnswer(t *testing.T) { + msg := new(dns.Msg) + msg.SetReply(new(dns.Msg)) + msg.Answer = append(msg.Answer, &dns.A{ + Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, + A: net.ParseIP("93.184.216.34"), + }) + + ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil) + r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil) + r.Process(msg) + + if r.Type != RespAnswer { + t.Errorf("Type = %d, want %d", r.Type, RespAnswer) + } + if r.Decoded == nil { + t.Fatal("Decoded should not be nil") + } +} + +func TestResponseClassifyReferral(t *testing.T) { + msg := new(dns.Msg) + msg.Rcode = dns.RcodeSuccess + msg.Authoritative = false + msg.Ns = append(msg.Ns, &dns.NS{ + Hdr: dns.RR_Header{Name: "com.", Rrtype: dns.TypeNS, Class: dns.ClassINET, Ttl: 172800}, + Ns: "a.gtld-servers.net.", + }) + msg.Extra = append(msg.Extra, &dns.A{ + Hdr: dns.RR_Header{Name: "a.gtld-servers.net.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 172800}, + A: net.ParseIP("192.5.6.30"), + }) + + ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil) + cache := NewInfoCache(nil) + r := NewResponse(ref, net.ParseIP("1.2.3.4"), cache) + r.Process(msg) + + if r.Type != RespReferral { + t.Errorf("Type = %d, want %d", r.Type, RespReferral) + } +} + +func TestResponseClassifyNXDOMAIN(t *testing.T) { + msg := new(dns.Msg) + msg.Rcode = dns.RcodeNameError + + ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil) + r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil) + r.Process(msg) + + if r.Type != RespNXDOMAIN { + t.Errorf("Type = %d, want %d", r.Type, RespNXDOMAIN) + } +} + +func TestResponseClassifyNODATA(t *testing.T) { + msg := new(dns.Msg) + msg.Rcode = dns.RcodeSuccess + msg.Authoritative = true + msg.Ns = append(msg.Ns, &dns.SOA{ + Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeSOA, Class: dns.ClassINET, Ttl: 3600}, + }) + + ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil) + r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil) + r.Process(msg) + + if r.Type != RespNODATA { + t.Errorf("Type = %d, want %d", r.Type, RespNODATA) + } +} + +func TestResponseClassifySERVFAIL(t *testing.T) { + msg := new(dns.Msg) + msg.Rcode = dns.RcodeServerFailure + + ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil) + r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil) + r.Process(msg) + + if r.Type != RespSERVFAIL { + t.Errorf("Type = %d, want %d", r.Type, RespSERVFAIL) + } +} + +func TestResponseCNAMEFollow(t *testing.T) { + msg := new(dns.Msg) + msg.SetReply(new(dns.Msg)) + msg.Answer = append(msg.Answer, + &dns.CNAME{ + Hdr: dns.RR_Header{Name: "www.example.com.", Rrtype: dns.TypeCNAME, Class: dns.ClassINET}, + Target: "example.com.", + }, + ) + + ref := NewReferral("www.example.com", dns.TypeA, ".", 0, 1.0, nil) + r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil) + r.Process(msg) + + if r.Type != RespCNAMEFollow { + t.Errorf("Type = %d, want %d", r.Type, RespCNAMEFollow) + } +} + +func TestResponseCNAMEWithFinalAnswer(t *testing.T) { + msg := new(dns.Msg) + msg.SetReply(new(dns.Msg)) + msg.Answer = append(msg.Answer, + &dns.CNAME{ + Hdr: dns.RR_Header{Name: "www.example.com.", Rrtype: dns.TypeCNAME, Class: dns.ClassINET}, + Target: "example.com.", + }, + &dns.A{ + Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, + A: net.ParseIP("93.184.216.34"), + }, + ) + + ref := NewReferral("www.example.com", dns.TypeA, ".", 0, 1.0, nil) + r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil) + r.Process(msg) + + if r.Type != RespAnswer { + t.Errorf("Type = %d, want %d (CNAME with final A answer)", r.Type, RespAnswer) + } +} + +func TestResponseChildReferrals(t *testing.T) { + msg := new(dns.Msg) + msg.Rcode = dns.RcodeSuccess + msg.Authoritative = false + msg.Ns = append(msg.Ns, + &dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeNS}, Ns: "a.gtld-servers.net."}, + &dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeNS}, Ns: "b.gtld-servers.net."}, + ) + msg.Extra = append(msg.Extra, + &dns.A{Hdr: dns.RR_Header{Name: "a.gtld-servers.net.", Rrtype: dns.TypeA}, A: net.ParseIP("192.5.6.30")}, + &dns.A{Hdr: dns.RR_Header{Name: "b.gtld-servers.net.", Rrtype: dns.TypeA}, A: net.ParseIP("192.33.14.30")}, + ) + + ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil) + cache := NewInfoCache(nil) + r := NewResponse(ref, net.ParseIP("198.41.0.4"), cache) + r.Process(msg) + + children := r.ChildReferrals() + if len(children) != 2 { + t.Fatalf("expected 2 child referrals, got %d", len(children)) + } + + if children[0].Name != "example.com." { + t.Errorf("child[0] name = %q, want example.com.", children[0].Name) + } + if children[0].Bailiwick != "a.gtld-servers.net." { + t.Errorf("child[0] bailiwick = %q, want a.gtld-servers.net.", children[0].Bailiwick) + } + if children[0].Prob != 0.5 { + t.Errorf("child[0] prob = %f, want 0.5", children[0].Prob) + } + if children[0].Depth != 1 { + t.Errorf("child[0] depth = %d, want 1", children[0].Depth) + } + if !children[0].HasAddresses() { + t.Error("child[0] should have glue addresses") + } + + if !children[1].HasAddresses() { + t.Error("child[1] should have glue addresses") + } + + nsNames := cache.LookupNS("example.com.") + if len(nsNames) != 2 { + t.Errorf("expected 2 NS in cache, got %d", len(nsNames)) + } +} + +func TestResponseChildReferralsNonReferral(t *testing.T) { + msg := new(dns.Msg) + msg.SetReply(new(dns.Msg)) + msg.Answer = append(msg.Answer, &dns.A{ + Hdr: dns.RR_Header{Rrtype: dns.TypeA}, + A: net.ParseIP("1.2.3.4"), + }) + + ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil) + r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil) + r.Process(msg) + + if children := r.ChildReferrals(); children != nil { + t.Error("non-referral should not produce child referrals") + } +} + +func TestResponseCNAMEFollowReferral(t *testing.T) { + msg := new(dns.Msg) + msg.SetReply(new(dns.Msg)) + msg.Answer = append(msg.Answer, + &dns.CNAME{ + Hdr: dns.RR_Header{Name: "www.example.com.", Rrtype: dns.TypeCNAME}, + Target: "example.com.", + }, + ) + + ref := NewReferral("www.example.com", dns.TypeA, ".", 0, 1.0, nil) + r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil) + r.Process(msg) + + follow := r.CNAMEFollowReferral() + if follow == nil { + t.Fatal("expected CNAME follow referral") + } + if follow.Name != "example.com." { + t.Errorf("follow name = %q, want %q", follow.Name, "example.com.") + } + if follow.Depth != 1 { + t.Errorf("follow depth = %d, want 1", follow.Depth) + } +} + +func TestResponseIsTerminal(t *testing.T) { + tests := []struct { + respType ResponseType + want bool + }{ + {RespAnswer, true}, + {RespNODATA, true}, + {RespNXDOMAIN, true}, + {RespSERVFAIL, true}, + {RespError, true}, + {RespReferral, false}, + {RespCNAMEFollow, false}, + } + + for _, tt := range tests { + t.Run(tt.respType.String(), func(t *testing.T) { + r := &Response{Type: tt.respType} + if got := r.IsTerminal(); got != tt.want { + t.Errorf("IsTerminal() = %v, want %v", got, tt.want) + } + }) + } +} + +func TestResponseTypeString(t *testing.T) { + tests := []struct { + rt ResponseType + want string + }{ + {RespReferral, "referral"}, + {RespAnswer, "answer"}, + {RespCNAMEFollow, "cname_follow"}, + {RespNODATA, "nodata"}, + {RespNXDOMAIN, "nxdomain"}, + {RespSERVFAIL, "servfail"}, + {RespError, "error"}, + } + + for _, tt := range tests { + t.Run(tt.want, func(t *testing.T) { + if got := tt.rt.String(); got != tt.want { + t.Errorf("String() = %q, want %q", got, tt.want) + } + }) + } +} + +func TestResponseChildReferralsProbabilityInheritance(t *testing.T) { + msg := new(dns.Msg) + msg.Rcode = dns.RcodeSuccess + msg.Authoritative = false + msg.Ns = append(msg.Ns, + &dns.NS{Hdr: dns.RR_Header{Name: "com.", Rrtype: dns.TypeNS}, Ns: "a.gtld-servers.net."}, + &dns.NS{Hdr: dns.RR_Header{Name: "com.", Rrtype: dns.TypeNS}, Ns: "b.gtld-servers.net."}, + &dns.NS{Hdr: dns.RR_Header{Name: "com.", Rrtype: dns.TypeNS}, Ns: "c.gtld-servers.net."}, + ) + + ref := NewReferral("example.com", dns.TypeA, ".", 0, 0.5, nil) + r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil) + r.Process(msg) + + children := r.ChildReferrals() + if len(children) != 3 { + t.Fatalf("expected 3 children, got %d", len(children)) + } + + for _, c := range children { + if c.Prob != 0.5/3.0 { + t.Errorf("child prob = %f, want %f", c.Prob, 0.5/3.0) + } + } +} + +func TestResponseChildReferralsEmptyAuthority(t *testing.T) { + msg := new(dns.Msg) + msg.Rcode = dns.RcodeSuccess + msg.Authoritative = false + + ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil) + r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil) + r.Process(msg) + + children := r.ChildReferrals() + if len(children) != 0 { + t.Errorf("expected 0 children with empty authority, got %d", len(children)) + } +} + +func TestResponseNilReferral(t *testing.T) { + r := NewResponse(nil, net.ParseIP("1.2.3.4"), nil) + children := r.ChildReferrals() + if children != nil { + t.Error("nil referral should produce no children") + } + + follow := r.CNAMEFollowReferral() + if follow != nil { + t.Error("nil referral should produce no CNAME follow") + } +} diff --git a/internal/traverse/stack.go b/internal/traverse/stack.go new file mode 100644 index 0000000..ecb7a36 --- /dev/null +++ b/internal/traverse/stack.go @@ -0,0 +1,58 @@ +package traverse + +const DefaultMaxDepth = 20 + +type Stack struct { + items []*Referral + maxDepth int +} + +func NewStack(maxDepth int) *Stack { + if maxDepth <= 0 { + maxDepth = DefaultMaxDepth + } + return &Stack{ + items: make([]*Referral, 0), + maxDepth: maxDepth, + } +} + +func (s *Stack) Push(r *Referral) bool { + if r == nil { + return false + } + if r.Depth >= s.maxDepth { + return false + } + s.items = append(s.items, r) + return true +} + +func (s *Stack) Pop() *Referral { + if len(s.items) == 0 { + return nil + } + idx := len(s.items) - 1 + item := s.items[idx] + s.items = s.items[:idx] + return item +} + +func (s *Stack) Peek() *Referral { + if len(s.items) == 0 { + return nil + } + return s.items[len(s.items)-1] +} + +func (s *Stack) Len() int { + return len(s.items) +} + +func (s *Stack) MaxDepth() int { + return s.maxDepth +} + +func (s *Stack) IsEmpty() bool { + return len(s.items) == 0 +} diff --git a/internal/traverse/stack_test.go b/internal/traverse/stack_test.go new file mode 100644 index 0000000..3e663c6 --- /dev/null +++ b/internal/traverse/stack_test.go @@ -0,0 +1,143 @@ +package traverse + +import ( + "testing" + + "github.com/hits/ExploreDNS/internal/dns" +) + +func TestNewStack(t *testing.T) { + s := NewStack(10) + if s.MaxDepth() != 10 { + t.Errorf("MaxDepth = %d, want 10", s.MaxDepth()) + } + if !s.IsEmpty() { + t.Error("new stack should be empty") + } + if s.Len() != 0 { + t.Errorf("Len = %d, want 0", s.Len()) + } +} + +func TestNewStackDefaultDepth(t *testing.T) { + s := NewStack(0) + if s.MaxDepth() != DefaultMaxDepth { + t.Errorf("MaxDepth = %d, want %d", s.MaxDepth(), DefaultMaxDepth) + } + + s = NewStack(-5) + if s.MaxDepth() != DefaultMaxDepth { + t.Errorf("MaxDepth = %d, want %d", s.MaxDepth(), DefaultMaxDepth) + } +} + +func TestStackPushPop(t *testing.T) { + s := NewStack(5) + ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil) + + ok := s.Push(ref) + if !ok { + t.Error("Push should succeed") + } + if s.Len() != 1 { + t.Errorf("Len = %d, want 1", s.Len()) + } + + popped := s.Pop() + if popped != ref { + t.Error("popped referral should match pushed") + } + if s.Len() != 0 { + t.Errorf("Len = %d, want 0", s.Len()) + } +} + +func TestStackLIFO(t *testing.T) { + s := NewStack(5) + r1 := NewReferral("a.com", dns.TypeA, ".", 0, 1.0, nil) + r2 := NewReferral("b.com", dns.TypeA, ".", 1, 1.0, nil) + r3 := NewReferral("c.com", dns.TypeA, ".", 2, 1.0, nil) + + s.Push(r1) + s.Push(r2) + s.Push(r3) + + if popped := s.Pop(); popped != r3 { + t.Error("should pop r3 first (LIFO)") + } + if popped := s.Pop(); popped != r2 { + t.Error("should pop r2 second") + } + if popped := s.Pop(); popped != r1 { + t.Error("should pop r1 third") + } +} + +func TestStackPushNil(t *testing.T) { + s := NewStack(5) + ok := s.Push(nil) + if ok { + t.Error("Push(nil) should return false") + } + if s.Len() != 0 { + t.Errorf("Len = %d, want 0", s.Len()) + } +} + +func TestStackMaxDepth(t *testing.T) { + s := NewStack(3) + + r0 := NewReferral("a.com", dns.TypeA, ".", 0, 1.0, nil) + r1 := NewReferral("b.com", dns.TypeA, ".", 1, 1.0, nil) + r2 := NewReferral("c.com", dns.TypeA, ".", 2, 1.0, nil) + r3 := NewReferral("d.com", dns.TypeA, ".", 3, 1.0, nil) + + if !s.Push(r0) { + t.Error("depth 0 should be accepted") + } + if !s.Push(r1) { + t.Error("depth 1 should be accepted") + } + if !s.Push(r2) { + t.Error("depth 2 should be accepted") + } + if s.Push(r3) { + t.Error("depth 3 should be rejected (maxDepth=3)") + } +} + +func TestStackPopEmpty(t *testing.T) { + s := NewStack(5) + if popped := s.Pop(); popped != nil { + t.Error("Pop on empty stack should return nil") + } +} + +func TestStackPeek(t *testing.T) { + s := NewStack(5) + if peek := s.Peek(); peek != nil { + t.Error("Peek on empty stack should return nil") + } + + ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil) + s.Push(ref) + + if peek := s.Peek(); peek != ref { + t.Error("Peek should return top item") + } + if s.Len() != 1 { + t.Errorf("Peek should not remove item, Len = %d, want 1", s.Len()) + } +} + +func TestStackIsEmpty(t *testing.T) { + s := NewStack(5) + if !s.IsEmpty() { + t.Error("new stack should be empty") + } + + s.Push(NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil)) + if s.IsEmpty() { + t.Error("stack with item should not be empty") + } +} diff --git a/internal/traverse/traverser.go b/internal/traverse/traverser.go new file mode 100644 index 0000000..f98078d --- /dev/null +++ b/internal/traverse/traverser.go @@ -0,0 +1,274 @@ +package traverse + +import ( + "context" + "fmt" + "net" + "sync" + + "github.com/hits/ExploreDNS/internal/dns" + miekgdns "github.com/miekg/dns" +) + +type TraverserConfig struct { + MaxDepth int + QueryType uint16 + RootConfig *dns.RootDiscoveryConfig + QueryConfig *dns.QueryConfig + RootAddrs []net.IP +} + +func DefaultTraverserConfig() *TraverserConfig { + return &TraverserConfig{ + MaxDepth: DefaultMaxDepth, + QueryType: dns.TypeA, + RootConfig: nil, + QueryConfig: nil, + RootAddrs: nil, + } +} + +type TraversalResult struct { + Referral *Referral + Response *Response +} + +type Traverser struct { + config *TraverserConfig + exchange dns.ExchangeFunc +} + +func NewTraverser(cfg *TraverserConfig) *Traverser { + if cfg == nil { + cfg = DefaultTraverserConfig() + } + return &Traverser{ + config: cfg, + exchange: nil, + } +} + +func (t *Traverser) SetExchange(fn dns.ExchangeFunc) { + t.exchange = fn +} + +func (t *Traverser) Traverse(ctx context.Context, name string) ([]TraversalResult, error) { + name = miekgdns.Fqdn(name) + + roots, err := t.discoverRoots(ctx) + if err != nil { + return nil, fmt.Errorf("root discovery: %w", err) + } + + initial := NewReferral(name, t.config.QueryType, ".", 0, 1.0, nil) + initial.Addresses = roots + + stack := NewStack(t.config.MaxDepth) + stack.Push(initial) + + rootCache := NewInfoCache(nil) + var ( + mu sync.Mutex + results []TraversalResult + ) + + for { + select { + case <-ctx.Done(): + return results, fmt.Errorf("traversal cancelled: %w", ctx.Err()) + default: + } + + ref := stack.Pop() + if ref == nil { + break + } + + cache := rootCache + if ref.Parent != nil { + cache = rootCache.Child() + } + + resp := t.processReferral(ctx, ref, cache) + mu.Lock() + results = append(results, TraversalResult{Referral: ref, Response: resp}) + mu.Unlock() + + if resp.IsTerminal() { + continue + } + + if resp.Type == RespReferral { + children := resp.ChildReferrals() + for _, child := range children { + if !stack.Push(child) { + mu.Lock() + results = append(results, TraversalResult{ + Referral: child, + Response: &Response{ + Referral: child, + Type: RespError, + }, + }) + mu.Unlock() + } + } + } + + if resp.Type == RespCNAMEFollow { + follow := resp.CNAMEFollowReferral() + if follow != nil { + if !stack.Push(follow) { + mu.Lock() + results = append(results, TraversalResult{ + Referral: follow, + Response: &Response{ + Referral: follow, + Type: RespError, + }, + }) + mu.Unlock() + } + } + } + } + + return results, nil +} + +func (t *Traverser) discoverRoots(ctx context.Context) ([]net.IP, error) { + if len(t.config.RootAddrs) > 0 { + return t.config.RootAddrs, nil + } + + servers, err := dns.DiscoverRoots(ctx, t.config.RootConfig) + if err != nil { + return nil, err + } + + var addrs []net.IP + for _, srv := range servers { + addrs = append(addrs, srv.AllIPs(false)...) + } + return addrs, nil +} + +func (t *Traverser) processReferral(ctx context.Context, ref *Referral, cache *InfoCache) *Response { + if !ref.HasAddresses() { + ref.Addresses = t.resolveGlueViaSystem(ctx, ref.Name, cache) + if len(ref.Addresses) > 0 { + ref.State = StateResolved + } else { + return &Response{ + Referral: ref, + Type: RespError, + } + } + } + + for _, addr := range ref.Addresses { + resp := t.queryServer(ctx, ref, addr, cache) + if resp != nil && resp.Type != RespSERVFAIL { + return resp + } + } + + return &Response{ + Referral: ref, + Type: RespSERVFAIL, + } +} + +func (t *Traverser) queryServer(ctx context.Context, ref *Referral, server net.IP, cache *InfoCache) *Response { + var msg *miekgdns.Msg + var err error + + if t.exchange != nil { + msg, err = t.iterativeQueryWithExchange(ctx, server, ref.Name, ref.Qtype) + } else { + msg, err = dns.Query(ctx, server, ref.Name, ref.Qtype, t.config.QueryConfig) + if err == nil { + msg = t.ensureRDFalse(msg, server, ref.Name, ref.Qtype, t.config.QueryConfig) + } + } + + if err != nil { + return &Response{ + Referral: ref, + Server: server, + Type: RespError, + } + } + + resp := NewResponse(ref, server, cache) + resp.Process(msg) + return resp +} + +func (t *Traverser) iterativeQueryWithExchange(ctx context.Context, server net.IP, name string, qtype uint16) (*miekgdns.Msg, error) { + if t.config.QueryConfig == nil { + return dns.IterativeQueryWithExchange(ctx, server, name, qtype, nil, t.exchange) + } + return dns.IterativeQueryWithExchange(ctx, server, name, qtype, t.config.QueryConfig, t.exchange) +} + +func (t *Traverser) ensureRDFalse(msg *miekgdns.Msg, server net.IP, name string, qtype uint16, cfg *dns.QueryConfig) *miekgdns.Msg { + if msg != nil && msg.RecursionDesired { + if t.exchange != nil { + ctx := context.Background() + var err error + msg, err = t.iterativeQueryWithExchange(ctx, server, name, qtype) + if err != nil { + return nil + } + return msg + } + msg.RecursionDesired = false + } + return msg +} + +func (t *Traverser) resolveGlueViaSystem(ctx context.Context, name string, cache *InfoCache) []net.IP { + if cache != nil { + if addrs := cache.LookupGlue(name); len(addrs) > 0 { + return addrs + } + } + + c := &miekgdns.Client{ + Net: "udp", + ReadTimeout: 5, + WriteTimeout: 5, + } + if deadline, ok := ctx.Deadline(); ok { + c.ReadTimeout = deadline.Sub(deadline) + c.WriteTimeout = deadline.Sub(deadline) + } + + fqdn := miekgdns.Fqdn(name) + + aMsg, _, err := c.ExchangeContext(ctx, newAQuery(fqdn), "127.0.0.1:53") + if err == nil { + var addrs []net.IP + for _, rr := range aMsg.Answer { + if a, ok := rr.(*miekgdns.A); ok { + addrs = append(addrs, a.A) + } + } + if len(addrs) > 0 { + if cache != nil { + cache.StoreGlue(name, addrs) + } + return addrs + } + } + + return nil +} + +func newAQuery(name string) *miekgdns.Msg { + m := new(miekgdns.Msg) + m.SetQuestion(name, miekgdns.TypeA) + m.RecursionDesired = true + return m +} diff --git a/internal/traverse/traverser_test.go b/internal/traverse/traverser_test.go new file mode 100644 index 0000000..cd41625 --- /dev/null +++ b/internal/traverse/traverser_test.go @@ -0,0 +1,483 @@ +package traverse + +import ( + "context" + "net" + "testing" + + "github.com/miekg/dns" +) + +const ( + dnsTypeA = dns.TypeA + dnsTypeNS = dns.TypeNS + dnsTypeCNAME = dns.TypeCNAME + dnsTypeSOA = dns.TypeSOA +) + +func TestDefaultTraverserConfig(t *testing.T) { + cfg := DefaultTraverserConfig() + if cfg.MaxDepth != DefaultMaxDepth { + t.Errorf("MaxDepth = %d, want %d", cfg.MaxDepth, DefaultMaxDepth) + } + if cfg.QueryType != dnsTypeA { + t.Errorf("QueryType = %d, want %d", cfg.QueryType, dnsTypeA) + } +} + +func TestNewTraverserNilConfig(t *testing.T) { + tr := NewTraverser(nil) + if tr == nil { + t.Fatal("NewTraverser(nil) should not return nil") + } +} + +func TestTraverserSimpleTraversal(t *testing.T) { + answerResp := func() *dns.Msg { + m := new(dns.Msg) + m.SetReply(new(dns.Msg)) + m.Answer = append(m.Answer, &dns.A{ + Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, + A: net.ParseIP("93.184.216.34"), + }) + return m + }() + + tr := NewTraverser(&TraverserConfig{ + MaxDepth: 5, + QueryType: dnsTypeA, + RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, + }) + tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + return answerResp.Copy(), nil + }) + + ctx := context.Background() + results, err := tr.Traverse(ctx, "example.com") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if len(results) == 0 { + t.Fatal("expected at least 1 result") + } + + found := false + for _, r := range results { + if r.Response.Type == RespAnswer { + found = true + break + } + } + if !found { + t.Error("expected to find an answer response") + } +} + +func TestTraverserReferralTraversal(t *testing.T) { + rootAnswer := new(dns.Msg) + rootAnswer.Rcode = dns.RcodeSuccess + rootAnswer.Authoritative = false + rootAnswer.Ns = append(rootAnswer.Ns, &dns.NS{ + Hdr: dns.RR_Header{Name: "com.", Rrtype: dnsTypeNS, Class: dns.ClassINET}, + Ns: "a.gtld-servers.net.", + }) + rootAnswer.Extra = append(rootAnswer.Extra, &dns.A{ + Hdr: dns.RR_Header{Name: "a.gtld-servers.net.", Rrtype: dnsTypeA}, + A: net.ParseIP("192.5.6.30"), + }) + + tldAnswer := new(dns.Msg) + tldAnswer.SetReply(new(dns.Msg)) + tldAnswer.Answer = append(tldAnswer.Answer, &dns.A{ + Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, + A: net.ParseIP("93.184.216.34"), + }) + + tr := NewTraverser(&TraverserConfig{ + MaxDepth: 5, + QueryType: dnsTypeA, + RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, + }) + tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + q := msg.Question[0] + key := q.Name + "/" + dns.TypeToString[q.Qtype] + if q.Name == "example.com." && server == "198.41.0.4" { + return rootAnswer.Copy(), nil + } + if q.Name == "example.com." { + return tldAnswer.Copy(), nil + } + _ = key + return nil, nil + }) + + ctx := context.Background() + results, err := tr.Traverse(ctx, "example.com") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if len(results) < 2 { + t.Fatalf("expected at least 2 results (referral + answer), got %d", len(results)) + } +} + +func TestTraverserMaxDepth(t *testing.T) { + callCount := 0 + tr := NewTraverser(&TraverserConfig{ + MaxDepth: 2, + QueryType: dnsTypeA, + RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, + }) + tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + callCount++ + m := new(dns.Msg) + m.Rcode = dns.RcodeSuccess + m.Authoritative = false + m.Ns = append(m.Ns, &dns.NS{ + Hdr: dns.RR_Header{Rrtype: dnsTypeNS}, + Ns: "ns.example.com.", + }) + m.Extra = append(m.Extra, &dns.A{ + Hdr: dns.RR_Header{Name: "ns.example.com.", Rrtype: dnsTypeA}, + A: net.ParseIP("1.2.3.4"), + }) + return m, nil + }) + + ctx := context.Background() + results, err := tr.Traverse(ctx, "deep.example.com") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + if callCount < 1 { + t.Errorf("expected at least 1 call before max depth, got %d", callCount) + } + + depthExceeded := false + for _, r := range results { + if r.Referral != nil && r.Referral.Depth >= 2 { + depthExceeded = true + } + if r.Response.Type == RespError { + depthExceeded = true + } + } + if !depthExceeded { + t.Error("expected to see depth exceeded results") + } +} + +func TestTraverserContextCancellation(t *testing.T) { + tr := NewTraverser(&TraverserConfig{ + MaxDepth: 5, + QueryType: dnsTypeA, + RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, + }) + tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + m := new(dns.Msg) + m.Rcode = dns.RcodeSuccess + m.Authoritative = false + m.Ns = append(m.Ns, &dns.NS{ + Hdr: dns.RR_Header{Rrtype: dnsTypeNS}, + Ns: "ns.example.com.", + }) + m.Extra = append(m.Extra, &dns.A{ + Hdr: dns.RR_Header{Name: "ns.example.com.", Rrtype: dnsTypeA}, + A: net.ParseIP("1.2.3.4"), + }) + return m, nil + }) + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + _, err := tr.Traverse(ctx, "example.com") + if err == nil { + t.Fatal("expected error on cancelled context") + } +} + +func TestTraverserNXDOMAIN(t *testing.T) { + nxdResp := new(dns.Msg) + nxdResp.Rcode = dns.RcodeNameError + + tr := NewTraverser(&TraverserConfig{ + MaxDepth: 5, + QueryType: dnsTypeA, + RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, + }) + tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + return nxdResp.Copy(), nil + }) + + ctx := context.Background() + results, err := tr.Traverse(ctx, "nonexistent.invalid") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if len(results) == 0 { + t.Fatal("expected at least 1 result") + } + if results[0].Response.Type != RespNXDOMAIN { + t.Errorf("Type = %d, want %d", results[0].Response.Type, RespNXDOMAIN) + } +} + +func TestTraverserSERVFAIL(t *testing.T) { + sfResp := new(dns.Msg) + sfResp.Rcode = dns.RcodeServerFailure + + tr := NewTraverser(&TraverserConfig{ + MaxDepth: 5, + QueryType: dnsTypeA, + RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, + }) + tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + return sfResp.Copy(), nil + }) + + ctx := context.Background() + results, err := tr.Traverse(ctx, "example.com") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if len(results) == 0 { + t.Fatal("expected at least 1 result") + } + if results[0].Response.Type != RespSERVFAIL { + t.Errorf("Type = %d, want %d", results[0].Response.Type, RespSERVFAIL) + } +} + +func TestTraverserCNAMEFollow(t *testing.T) { + cnameResp := new(dns.Msg) + cnameResp.SetReply(new(dns.Msg)) + cnameResp.Answer = append(cnameResp.Answer, + &dns.CNAME{ + Hdr: dns.RR_Header{Name: "www.example.com.", Rrtype: dnsTypeCNAME, Class: dns.ClassINET}, + Target: "example.com.", + }, + ) + + answerResp := new(dns.Msg) + answerResp.SetReply(new(dns.Msg)) + answerResp.Answer = append(answerResp.Answer, &dns.A{ + Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, + A: net.ParseIP("93.184.216.34"), + }) + + tr := NewTraverser(&TraverserConfig{ + MaxDepth: 5, + QueryType: dnsTypeA, + RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, + }) + tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + q := msg.Question[0] + if q.Name == "www.example.com." { + return cnameResp.Copy(), nil + } + if q.Name == "example.com." { + return answerResp.Copy(), nil + } + return nil, nil + }) + + ctx := context.Background() + results, err := tr.Traverse(ctx, "www.example.com") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + foundCNAME := false + foundAnswer := false + for _, r := range results { + if r.Response.Type == RespCNAMEFollow { + foundCNAME = true + } + if r.Response.Type == RespAnswer { + foundAnswer = true + } + } + if !foundCNAME { + t.Error("expected CNAME follow response") + } + if !foundAnswer { + t.Error("expected final answer response") + } +} + +func TestTraverserProbabilityCalculation(t *testing.T) { + rootReferral := new(dns.Msg) + rootReferral.Rcode = dns.RcodeSuccess + rootReferral.Authoritative = false + rootReferral.Ns = append(rootReferral.Ns, + &dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dnsTypeNS}, Ns: "a.root-servers.net."}, + &dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dnsTypeNS}, Ns: "b.root-servers.net."}, + ) + rootReferral.Extra = append(rootReferral.Extra, + &dns.A{Hdr: dns.RR_Header{Name: "a.root-servers.net.", Rrtype: dnsTypeA}, A: net.ParseIP("198.41.0.4")}, + &dns.A{Hdr: dns.RR_Header{Name: "b.root-servers.net.", Rrtype: dnsTypeA}, A: net.ParseIP("199.9.14.201")}, + ) + + tldAnswer := new(dns.Msg) + tldAnswer.SetReply(new(dns.Msg)) + tldAnswer.Answer = append(tldAnswer.Answer, &dns.A{ + Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, + A: net.ParseIP("93.184.216.34"), + }) + + tr := NewTraverser(&TraverserConfig{ + MaxDepth: 5, + QueryType: dnsTypeA, + RootAddrs: []net.IP{net.ParseIP("1.2.3.4")}, + }) + tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + q := msg.Question[0] + if q.Name == "example.com." && server == "1.2.3.4" { + return rootReferral.Copy(), nil + } + return tldAnswer.Copy(), nil + }) + + ctx := context.Background() + results, err := tr.Traverse(ctx, "example.com") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + for _, r := range results { + if r.Referral != nil && r.Referral.Depth == 1 && r.Referral.Parent != nil { + if r.Referral.Prob != 0.5 { + t.Errorf("child prob = %f, want 0.5", r.Referral.Prob) + } + } + } +} + +func TestTraverserNODATA(t *testing.T) { + nodataResp := new(dns.Msg) + nodataResp.Rcode = dns.RcodeSuccess + nodataResp.Authoritative = true + nodataResp.Ns = append(nodataResp.Ns, &dns.SOA{ + Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeSOA, Class: dns.ClassINET, Ttl: 3600}, + }) + + tr := NewTraverser(&TraverserConfig{ + MaxDepth: 5, + QueryType: dnsTypeA, + RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, + }) + tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + return nodataResp.Copy(), nil + }) + + ctx := context.Background() + results, err := tr.Traverse(ctx, "example.com") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if len(results) == 0 { + t.Fatal("expected at least 1 result") + } + if results[0].Response.Type != RespNODATA { + t.Errorf("Type = %d, want %d", results[0].Response.Type, RespNODATA) + } +} + +func TestTraverserMultipleRoots(t *testing.T) { + answerResp := new(dns.Msg) + answerResp.SetReply(new(dns.Msg)) + answerResp.Answer = append(answerResp.Answer, &dns.A{ + Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, + A: net.ParseIP("93.184.216.34"), + }) + + tr := NewTraverser(&TraverserConfig{ + MaxDepth: 5, + QueryType: dnsTypeA, + RootAddrs: []net.IP{net.ParseIP("198.41.0.4"), net.ParseIP("199.9.14.201")}, + }) + tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + return answerResp.Copy(), nil + }) + + ctx := context.Background() + results, err := tr.Traverse(ctx, "example.com") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if len(results) == 0 { + t.Fatal("expected results") + } +} + +func TestTraverserNilExchangeResponse(t *testing.T) { + tr := NewTraverser(&TraverserConfig{ + MaxDepth: 5, + QueryType: dnsTypeA, + RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, + }) + tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + return nil, nil + }) + + ctx := context.Background() + results, err := tr.Traverse(ctx, "example.com") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if len(results) == 0 { + t.Fatal("expected at least 1 result even with nil response") + } +} + +func TestTraverserCacheChaining(t *testing.T) { + rootReferral := new(dns.Msg) + rootReferral.Rcode = dns.RcodeSuccess + rootReferral.Authoritative = false + rootReferral.Ns = append(rootReferral.Ns, + &dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dnsTypeNS}, Ns: "a.gtld-servers.net."}, + ) + rootReferral.Extra = append(rootReferral.Extra, + &dns.A{Hdr: dns.RR_Header{Name: "a.gtld-servers.net.", Rrtype: dnsTypeA}, A: net.ParseIP("192.5.6.30")}, + ) + + answerResp := new(dns.Msg) + answerResp.SetReply(new(dns.Msg)) + answerResp.Answer = append(answerResp.Answer, &dns.A{ + Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, + A: net.ParseIP("93.184.216.34"), + }) + + tr := NewTraverser(&TraverserConfig{ + MaxDepth: 5, + QueryType: dnsTypeA, + RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, + }) + tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + q := msg.Question[0] + if q.Name == "example.com." && server == "198.41.0.4" { + return rootReferral.Copy(), nil + } + return answerResp.Copy(), nil + }) + + ctx := context.Background() + results, err := tr.Traverse(ctx, "example.com") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + cacheHits := 0 + for _, r := range results { + if r.Response != nil && r.Response.Cache != nil { + if r.Response.Cache.NSCount() > 0 { + cacheHits++ + } + } + } + if cacheHits == 0 { + t.Error("expected cache to store NS records from referrals") + } +}