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) }, } 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 }) _, err := tr.Traverse(context.Background(), "example.com") if err != nil { t.Fatalf("Traverse: %v", err) } if len(events) < 2 { t.Fatalf("expected start and complete events, got %d", len(events)) } 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) } }