package dns import ( "net" "testing" "github.com/miekg/dns" ) func newTestMsg(rcode int) *dns.Msg { m := new(dns.Msg) m.Rcode = rcode return m } func TestDecodeResponseNil(t *testing.T) { d := DecodeResponse(nil) if d != nil { t.Error("DecodeResponse(nil) should return nil") } } func TestDecodeResponseNXDOMAIN(t *testing.T) { msg := newTestMsg(dns.RcodeNameError) d := DecodeResponse(msg) if d.Classification != ResponseNXDOMAIN { t.Errorf("classification = %v, want NXDOMAIN", d.Classification) } if d.RcodeName != "NXDOMAIN" { t.Errorf("RcodeName = %q, want %q", d.RcodeName, "NXDOMAIN") } } func TestDecodeResponseSERVFAIL(t *testing.T) { msg := newTestMsg(dns.RcodeServerFailure) d := DecodeResponse(msg) if d.Classification != ResponseSERVFAIL { t.Errorf("classification = %v, want SERVFAIL", d.Classification) } } func TestDecodeResponseAnswer(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: MustParseIP("93.184.216.34"), }) d := DecodeResponse(msg) if d.Classification != ResponseAnswer { t.Errorf("classification = %v, want Answer", d.Classification) } if len(d.Answers) != 1 { t.Errorf("Answers count = %d, want 1", len(d.Answers)) } } func TestDecodeResponseReferral(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.", }) d := DecodeResponse(msg) if d.Classification != ResponseReferral { t.Errorf("classification = %v, want Referral", d.Classification) } } func TestDecodeResponseNODATA(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}, }) d := DecodeResponse(msg) if d.Classification != ResponseNODATA { t.Errorf("classification = %v, want NODATA", d.Classification) } } func TestDecodeResponseTruncated(t *testing.T) { msg := new(dns.Msg) msg.Truncated = true d := DecodeResponse(msg) if !d.Truncated { t.Error("Truncated should be true") } } func TestIsTruncated(t *testing.T) { tests := []struct { name string msg *dns.Msg want bool }{ {"nil", nil, false}, {"not truncated", new(dns.Msg), false}, {"truncated", &dns.Msg{MsgHdr: dns.MsgHdr{Truncated: true}}, true}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { if got := IsTruncated(tt.msg); got != tt.want { t.Errorf("IsTruncated() = %v, want %v", got, tt.want) } }) } } func TestRcodeName(t *testing.T) { tests := []struct { name string msg *dns.Msg want string }{ {"nil", nil, "UNKNOWN"}, {"NOERROR", &dns.Msg{MsgHdr: dns.MsgHdr{Rcode: dns.RcodeSuccess}}, "NOERROR"}, {"NXDOMAIN", &dns.Msg{MsgHdr: dns.MsgHdr{Rcode: dns.RcodeNameError}}, "NXDOMAIN"}, {"SERVFAIL", &dns.Msg{MsgHdr: dns.MsgHdr{Rcode: dns.RcodeServerFailure}}, "SERVFAIL"}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { if got := RcodeName(tt.msg); got != tt.want { t.Errorf("RcodeName() = %q, want %q", got, tt.want) } }) } } func TestExtractAnswers(t *testing.T) { msg := 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: MustParseIP("93.184.216.34"), }) ans := ExtractAnswers(msg) if len(ans) != 1 { t.Errorf("expected 1 answer, got %d", len(ans)) } if ans := ExtractAnswers(nil); ans != nil { t.Error("ExtractAnswers(nil) should return nil") } } func TestExtractAuthority(t *testing.T) { msg := new(dns.Msg) 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.", }) auth := ExtractAuthority(msg) if len(auth) != 1 { t.Errorf("expected 1 authority record, got %d", len(auth)) } if auth := ExtractAuthority(nil); auth != nil { t.Error("ExtractAuthority(nil) should return nil") } } func TestIsReferral(t *testing.T) { tests := []struct { name string msg *dns.Msg want bool }{ {"nil", nil, false}, {"NXDOMAIN", &dns.Msg{MsgHdr: dns.MsgHdr{Rcode: dns.RcodeNameError}}, false}, {"has answers", func() *dns.Msg { m := new(dns.Msg) m.Answer = append(m.Answer, &dns.A{Hdr: dns.RR_Header{Rrtype: dns.TypeA}}) return m }(), false}, {"referral", func() *dns.Msg { m := new(dns.Msg) m.Rcode = dns.RcodeSuccess m.Authoritative = false m.Ns = append(m.Ns, &dns.NS{Hdr: dns.RR_Header{Rrtype: dns.TypeNS}}) return m }(), true}, {"authoritative NS is not referral", func() *dns.Msg { m := new(dns.Msg) m.Rcode = dns.RcodeSuccess m.Authoritative = true m.Ns = append(m.Ns, &dns.NS{Hdr: dns.RR_Header{Rrtype: dns.TypeNS}}) return m }(), false}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { if got := IsReferral(tt.msg); got != tt.want { t.Errorf("IsReferral() = %v, want %v", got, tt.want) } }) } } func TestIsNODATA(t *testing.T) { tests := []struct { name string msg *dns.Msg want bool }{ {"nil", nil, false}, {"non-success rcode", &dns.Msg{MsgHdr: dns.MsgHdr{Rcode: dns.RcodeNameError}}, false}, {"has answers", func() *dns.Msg { m := new(dns.Msg) m.Answer = append(m.Answer, &dns.A{Hdr: dns.RR_Header{Rrtype: dns.TypeA}}) return m }(), false}, {"referral is not NODATA", func() *dns.Msg { m := new(dns.Msg) m.Rcode = dns.RcodeSuccess m.Authoritative = false m.Ns = append(m.Ns, &dns.NS{Hdr: dns.RR_Header{Rrtype: dns.TypeNS}}) return m }(), false}, {"empty answer, authoritative", func() *dns.Msg { m := new(dns.Msg) m.Rcode = dns.RcodeSuccess m.Authoritative = true return m }(), true}, {"empty response, not authoritative", func() *dns.Msg { m := new(dns.Msg) m.Rcode = dns.RcodeSuccess return m }(), true}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { if got := IsNODATA(tt.msg); got != tt.want { t.Errorf("IsNODATA() = %v, want %v", got, tt.want) } }) } } func TestCNAMEChain(t *testing.T) { t.Run("nil msg", func(t *testing.T) { if chain := ExtractCNAMEChain(nil); chain != nil { t.Errorf("expected nil, got %v", chain) } if HasCNAMEChain(nil) { t.Error("HasCNAMEChain(nil) should be false") } }) t.Run("no CNAME records", func(t *testing.T) { msg := new(dns.Msg) msg.Answer = append(msg.Answer, &dns.A{ Hdr: dns.RR_Header{Rrtype: dns.TypeA}, }) if chain := ExtractCNAMEChain(msg); len(chain) != 0 { t.Errorf("expected empty chain, got %v", chain) } }) t.Run("CNAME chain", func(t *testing.T) { msg := new(dns.Msg) msg.Answer = append(msg.Answer, &dns.CNAME{ Hdr: dns.RR_Header{Name: "www.example.com.", Rrtype: dns.TypeCNAME}, Target: "example.com.", }, &dns.A{ Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA}, A: MustParseIP("93.184.216.34"), }, ) chain := ExtractCNAMEChain(msg) if len(chain) != 1 { t.Fatalf("expected 1 CNAME target, got %d", len(chain)) } if chain[0] != "example.com." { t.Errorf("CNAME target = %q, want %q", chain[0], "example.com.") } if !HasCNAMEChain(msg) { t.Error("HasCNAMEChain should be true") } }) t.Run("dedup CNAME targets", func(t *testing.T) { msg := new(dns.Msg) msg.Answer = append(msg.Answer, &dns.CNAME{ Hdr: dns.RR_Header{Name: "a.example.com.", Rrtype: dns.TypeCNAME}, Target: "b.example.com.", }, &dns.CNAME{ Hdr: dns.RR_Header{Name: "b.example.com.", Rrtype: dns.TypeCNAME}, Target: "b.example.com.", }, ) chain := ExtractCNAMEChain(msg) if len(chain) != 1 { t.Errorf("expected 1 deduped CNAME target, got %d: %v", len(chain), chain) } }) } func TestResponseClassificationString(t *testing.T) { tests := []struct { rc ResponseClassification want string }{ {ResponseAnswer, "answer"}, {ResponseReferral, "referral"}, {ResponseNODATA, "nodata"}, {ResponseNXDOMAIN, "nxdomain"}, {ResponseSERVFAIL, "servfail"}, {ResponseOther, "other"}, } for _, tt := range tests { t.Run(tt.want, func(t *testing.T) { if got := tt.rc.String(); got != tt.want { t.Errorf("String() = %q, want %q", got, tt.want) } }) } } func TestFormatRecord(t *testing.T) { t.Run("nil", func(t *testing.T) { if got := FormatRecord(nil); got != "" { t.Errorf("FormatRecord(nil) = %q, want empty", got) } }) t.Run("A record", func(t *testing.T) { rr := &dns.A{ Hdr: dns.RR_Header{ Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300, }, A: MustParseIP("93.184.216.34"), } got := FormatRecord(rr) if got == "" { t.Error("FormatRecord returned empty string") } }) } func TestDecodeResponseFlags(t *testing.T) { msg := new(dns.Msg) msg.RecursionAvailable = true msg.Authoritative = true d := DecodeResponse(msg) if !d.RecursionAvailable { t.Error("RecursionAvailable should be true") } if !d.Authoritative { t.Error("Authoritative should be true") } } func MustParseIP(s string) []byte { ip := net.ParseIP(s) if ip == nil { panic("invalid IP: " + s) } return ip.To4() }