package traverse import ( "context" "errors" "net" "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") } } }