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