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) }