Phase 1.2 implementation: - internal/dns/types.go: DNS record type constants (A, AAAA, NS, CNAME, SOA, MX, TXT, SRV, PTR, ANY) with QNameType helper and EDNS0 defaults - internal/dns/query.go: Query/QueryWithExchange with UDP, configurable EDNS0 buffer size, TCP fallback on truncation, always-TCP mode, configurable retries with timeout, context support, and injectable ExchangeFunc for testing - internal/dns/decode.go: Response classification (answer, referral, NODATA, NXDOMAIN, SERVFAIL), CNAME chain extraction with dedup, truncation/RCODE detection, and section extraction utilities 32 unit tests covering all acceptance criteria. Co-authored-by: multica-agent <github@multica.ai>
382 lines
9.2 KiB
Go
382 lines
9.2 KiB
Go
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()
|
|
}
|