Merge pull request 'feat: add DNS query/response layer with miekg/dns' (#3) from agent/go-expert-developer/ad0c58a4 into main
CI / test (push) Has been cancelled
CI / test (push) Has been cancelled
Reviewed-on: http://gitea.hansenits.com.au/hits/ExploreDNS/pulls/3
This commit was merged in pull request #3.
This commit is contained in:
@@ -2,8 +2,9 @@ module github.com/hits/ExploreDNS
|
|||||||
|
|
||||||
go 1.25.6
|
go 1.25.6
|
||||||
|
|
||||||
|
require github.com/miekg/dns v1.1.72
|
||||||
|
|
||||||
require (
|
require (
|
||||||
github.com/miekg/dns v1.1.72 // indirect
|
|
||||||
golang.org/x/mod v0.31.0 // indirect
|
golang.org/x/mod v0.31.0 // indirect
|
||||||
golang.org/x/net v0.48.0 // indirect
|
golang.org/x/net v0.48.0 // indirect
|
||||||
golang.org/x/sync v0.19.0 // indirect
|
golang.org/x/sync v0.19.0 // indirect
|
||||||
|
|||||||
@@ -1,3 +1,5 @@
|
|||||||
|
github.com/google/go-cmp v0.6.0 h1:ofyhxvXcZhMsU5ulbFiLKl/XBFqE1GSq7atu8tAmTRI=
|
||||||
|
github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY=
|
||||||
github.com/miekg/dns v1.1.72 h1:vhmr+TF2A3tuoGNkLDFK9zi36F2LS+hKTRW0Uf8kbzI=
|
github.com/miekg/dns v1.1.72 h1:vhmr+TF2A3tuoGNkLDFK9zi36F2LS+hKTRW0Uf8kbzI=
|
||||||
github.com/miekg/dns v1.1.72/go.mod h1:+EuEPhdHOsfk6Wk5TT2CzssZdqkmFhf8r+aVyDEToIs=
|
github.com/miekg/dns v1.1.72/go.mod h1:+EuEPhdHOsfk6Wk5TT2CzssZdqkmFhf8r+aVyDEToIs=
|
||||||
golang.org/x/mod v0.31.0 h1:HaW9xtz0+kOcWKwli0ZXy79Ix+UW/vOfmWI5QVd2tgI=
|
golang.org/x/mod v0.31.0 h1:HaW9xtz0+kOcWKwli0ZXy79Ix+UW/vOfmWI5QVd2tgI=
|
||||||
|
|||||||
@@ -0,0 +1,199 @@
|
|||||||
|
package dns
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
|
||||||
|
"github.com/miekg/dns"
|
||||||
|
)
|
||||||
|
|
||||||
|
type ResponseClassification int
|
||||||
|
|
||||||
|
const (
|
||||||
|
ResponseAnswer ResponseClassification = iota
|
||||||
|
ResponseReferral
|
||||||
|
ResponseNODATA
|
||||||
|
ResponseNXDOMAIN
|
||||||
|
ResponseSERVFAIL
|
||||||
|
ResponseOther
|
||||||
|
)
|
||||||
|
|
||||||
|
func (rc ResponseClassification) String() string {
|
||||||
|
switch rc {
|
||||||
|
case ResponseAnswer:
|
||||||
|
return "answer"
|
||||||
|
case ResponseReferral:
|
||||||
|
return "referral"
|
||||||
|
case ResponseNODATA:
|
||||||
|
return "nodata"
|
||||||
|
case ResponseNXDOMAIN:
|
||||||
|
return "nxdomain"
|
||||||
|
case ResponseSERVFAIL:
|
||||||
|
return "servfail"
|
||||||
|
default:
|
||||||
|
return "other"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type DecodedResponse struct {
|
||||||
|
Rcode int
|
||||||
|
RcodeName string
|
||||||
|
Truncated bool
|
||||||
|
RecursionAvailable bool
|
||||||
|
Authoritative bool
|
||||||
|
Classification ResponseClassification
|
||||||
|
Answers []dns.RR
|
||||||
|
Authority []dns.RR
|
||||||
|
Additional []dns.RR
|
||||||
|
CNAMEChain []string
|
||||||
|
}
|
||||||
|
|
||||||
|
func DecodeResponse(msg *dns.Msg) *DecodedResponse {
|
||||||
|
if msg == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
d := &DecodedResponse{
|
||||||
|
Rcode: msg.Rcode,
|
||||||
|
RcodeName: dns.RcodeToString[msg.Rcode],
|
||||||
|
Truncated: msg.Truncated,
|
||||||
|
RecursionAvailable: msg.RecursionAvailable,
|
||||||
|
Authoritative: msg.Authoritative,
|
||||||
|
Answers: msg.Answer,
|
||||||
|
Authority: msg.Ns,
|
||||||
|
Additional: msg.Extra,
|
||||||
|
CNAMEChain: extractCNAMEChain(msg),
|
||||||
|
}
|
||||||
|
|
||||||
|
d.Classification = classify(msg)
|
||||||
|
|
||||||
|
return d
|
||||||
|
}
|
||||||
|
|
||||||
|
func classify(msg *dns.Msg) ResponseClassification {
|
||||||
|
switch msg.Rcode {
|
||||||
|
case dns.RcodeNameError:
|
||||||
|
return ResponseNXDOMAIN
|
||||||
|
case dns.RcodeServerFailure:
|
||||||
|
return ResponseSERVFAIL
|
||||||
|
case dns.RcodeSuccess:
|
||||||
|
return classifySuccess(msg)
|
||||||
|
default:
|
||||||
|
return ResponseOther
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func classifySuccess(msg *dns.Msg) ResponseClassification {
|
||||||
|
hasAnswers := len(msg.Answer) > 0
|
||||||
|
|
||||||
|
if hasAnswers {
|
||||||
|
return ResponseAnswer
|
||||||
|
}
|
||||||
|
|
||||||
|
hasNS := hasNSRecords(msg.Ns)
|
||||||
|
if hasNS && !msg.Authoritative {
|
||||||
|
return ResponseReferral
|
||||||
|
}
|
||||||
|
|
||||||
|
if hasNS {
|
||||||
|
return ResponseNODATA
|
||||||
|
}
|
||||||
|
|
||||||
|
return ResponseNODATA
|
||||||
|
}
|
||||||
|
|
||||||
|
func hasNSRecords(rrs []dns.RR) bool {
|
||||||
|
for _, rr := range rrs {
|
||||||
|
if _, ok := rr.(*dns.NS); ok {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
func extractCNAMEChain(msg *dns.Msg) []string {
|
||||||
|
var chain []string
|
||||||
|
seen := make(map[string]bool)
|
||||||
|
for _, rr := range msg.Answer {
|
||||||
|
if cname, ok := rr.(*dns.CNAME); ok {
|
||||||
|
target := cname.Target
|
||||||
|
if !seen[target] {
|
||||||
|
seen[target] = true
|
||||||
|
chain = append(chain, target)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return chain
|
||||||
|
}
|
||||||
|
|
||||||
|
func IsTruncated(msg *dns.Msg) bool {
|
||||||
|
return msg != nil && msg.Truncated
|
||||||
|
}
|
||||||
|
|
||||||
|
func RcodeName(msg *dns.Msg) string {
|
||||||
|
if msg == nil {
|
||||||
|
return "UNKNOWN"
|
||||||
|
}
|
||||||
|
return dns.RcodeToString[msg.Rcode]
|
||||||
|
}
|
||||||
|
|
||||||
|
func ExtractAnswers(msg *dns.Msg) []dns.RR {
|
||||||
|
if msg == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return msg.Answer
|
||||||
|
}
|
||||||
|
|
||||||
|
func ExtractAuthority(msg *dns.Msg) []dns.RR {
|
||||||
|
if msg == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return msg.Ns
|
||||||
|
}
|
||||||
|
|
||||||
|
func ExtractCNAMEChain(msg *dns.Msg) []string {
|
||||||
|
if msg == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return extractCNAMEChain(msg)
|
||||||
|
}
|
||||||
|
|
||||||
|
func IsReferral(msg *dns.Msg) bool {
|
||||||
|
if msg == nil || msg.Rcode != dns.RcodeSuccess || len(msg.Answer) > 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return hasNSRecords(msg.Ns) && !msg.Authoritative
|
||||||
|
}
|
||||||
|
|
||||||
|
func IsNODATA(msg *dns.Msg) bool {
|
||||||
|
if msg == nil || msg.Rcode != dns.RcodeSuccess {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if len(msg.Answer) > 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if IsReferral(msg) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func HasCNAMEChain(msg *dns.Msg) bool {
|
||||||
|
if msg == nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return len(extractCNAMEChain(msg)) > 0
|
||||||
|
}
|
||||||
|
|
||||||
|
func FormatRecord(rr dns.RR) string {
|
||||||
|
if rr == nil {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
header := rr.Header()
|
||||||
|
return fmt.Sprintf("%s %d %s %s %s",
|
||||||
|
header.Name,
|
||||||
|
header.Ttl,
|
||||||
|
dns.ClassToString[header.Class],
|
||||||
|
QNameType(header.Rrtype),
|
||||||
|
rr.String(),
|
||||||
|
)
|
||||||
|
}
|
||||||
@@ -0,0 +1,381 @@
|
|||||||
|
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()
|
||||||
|
}
|
||||||
@@ -0,0 +1,134 @@
|
|||||||
|
package dns
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"net"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/miekg/dns"
|
||||||
|
)
|
||||||
|
|
||||||
|
type QueryConfig struct {
|
||||||
|
UDPSize int
|
||||||
|
Timeout time.Duration
|
||||||
|
Retries int
|
||||||
|
UseTCP bool
|
||||||
|
}
|
||||||
|
|
||||||
|
func DefaultQueryConfig() *QueryConfig {
|
||||||
|
return &QueryConfig{
|
||||||
|
UDPSize: DefaultEDNS0UDPSize(),
|
||||||
|
Timeout: 5 * time.Second,
|
||||||
|
Retries: 3,
|
||||||
|
UseTCP: false,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type ExchangeFunc func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error)
|
||||||
|
|
||||||
|
func Query(ctx context.Context, server net.IP, name string, qtype uint16, cfg *QueryConfig) (*dns.Msg, error) {
|
||||||
|
if cfg == nil {
|
||||||
|
cfg = DefaultQueryConfig()
|
||||||
|
}
|
||||||
|
|
||||||
|
if cfg.UDPSize <= 0 {
|
||||||
|
cfg.UDPSize = DefaultEDNS0UDPSize()
|
||||||
|
}
|
||||||
|
if cfg.Timeout <= 0 {
|
||||||
|
cfg.Timeout = 5 * time.Second
|
||||||
|
}
|
||||||
|
|
||||||
|
return QueryWithExchange(ctx, server, name, qtype, cfg, realExchange)
|
||||||
|
}
|
||||||
|
|
||||||
|
func realExchange(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||||
|
addr := net.JoinHostPort(server, "53")
|
||||||
|
|
||||||
|
var c *dns.Client
|
||||||
|
if useTCP {
|
||||||
|
c = &dns.Client{
|
||||||
|
Net: "tcp",
|
||||||
|
ReadTimeout: 5 * time.Second,
|
||||||
|
WriteTimeout: 5 * time.Second,
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
c = &dns.Client{
|
||||||
|
Net: "udp",
|
||||||
|
ReadTimeout: 5 * time.Second,
|
||||||
|
WriteTimeout: 5 * time.Second,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if deadline, ok := ctx.Deadline(); ok {
|
||||||
|
c.ReadTimeout = time.Until(deadline)
|
||||||
|
c.WriteTimeout = time.Until(deadline)
|
||||||
|
}
|
||||||
|
|
||||||
|
r, _, err := c.ExchangeContext(ctx, msg, addr)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("dns exchange (%s) with %s: %w", c.Net, addr, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return r, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func QueryWithExchange(ctx context.Context, server net.IP, name string, qtype uint16, cfg *QueryConfig, exchangeFn ExchangeFunc) (*dns.Msg, error) {
|
||||||
|
if cfg == nil {
|
||||||
|
cfg = DefaultQueryConfig()
|
||||||
|
}
|
||||||
|
if cfg.UDPSize <= 0 {
|
||||||
|
cfg.UDPSize = DefaultEDNS0UDPSize()
|
||||||
|
}
|
||||||
|
|
||||||
|
msg := buildQuery(name, qtype, cfg.UDPSize)
|
||||||
|
serverStr := server.String()
|
||||||
|
|
||||||
|
var lastErr error
|
||||||
|
|
||||||
|
for attempt := 0; attempt < cfg.Retries; attempt++ {
|
||||||
|
if attempt > 0 {
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
return nil, fmt.Errorf("query retries cancelled: %w", ctx.Err())
|
||||||
|
case <-time.After(100 * time.Millisecond):
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if cfg.UseTCP {
|
||||||
|
resp, err := exchangeFn(ctx, serverStr, msg, true)
|
||||||
|
if err != nil {
|
||||||
|
lastErr = err
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
return resp, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
resp, err := exchangeFn(ctx, serverStr, msg, false)
|
||||||
|
if err != nil {
|
||||||
|
lastErr = err
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if resp.Truncated {
|
||||||
|
resp, err = exchangeFn(ctx, serverStr, msg, true)
|
||||||
|
if err != nil {
|
||||||
|
lastErr = err
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
return resp, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return resp, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil, fmt.Errorf("query %s %s failed after %d retries: %w", name, QNameType(qtype), cfg.Retries, lastErr)
|
||||||
|
}
|
||||||
|
|
||||||
|
func buildQuery(name string, qtype uint16, udpSize int) *dns.Msg {
|
||||||
|
m := new(dns.Msg)
|
||||||
|
m.SetQuestion(dns.Fqdn(name), qtype)
|
||||||
|
m.RecursionDesired = true
|
||||||
|
m.SetEdns0(uint16(udpSize), false)
|
||||||
|
return m
|
||||||
|
}
|
||||||
@@ -0,0 +1,289 @@
|
|||||||
|
package dns
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"net"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/miekg/dns"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestDefaultQueryConfig(t *testing.T) {
|
||||||
|
cfg := DefaultQueryConfig()
|
||||||
|
if cfg == nil {
|
||||||
|
t.Fatal("DefaultQueryConfig returned nil")
|
||||||
|
}
|
||||||
|
if cfg.UDPSize != 2048 {
|
||||||
|
t.Errorf("UDPSize = %d, want 2048", cfg.UDPSize)
|
||||||
|
}
|
||||||
|
if cfg.Retries != 3 {
|
||||||
|
t.Errorf("Retries = %d, want 3", cfg.Retries)
|
||||||
|
}
|
||||||
|
if cfg.UseTCP {
|
||||||
|
t.Error("UseTCP should be false by default")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildQuery(t *testing.T) {
|
||||||
|
msg := buildQuery("example.com.", TypeA, 2048)
|
||||||
|
if len(msg.Question) != 1 {
|
||||||
|
t.Fatalf("expected 1 question, got %d", len(msg.Question))
|
||||||
|
}
|
||||||
|
q := msg.Question[0]
|
||||||
|
if q.Name != "example.com." {
|
||||||
|
t.Errorf("question name = %q, want %q", q.Name, "example.com.")
|
||||||
|
}
|
||||||
|
if q.Qtype != TypeA {
|
||||||
|
t.Errorf("question type = %d, want %d", q.Qtype, TypeA)
|
||||||
|
}
|
||||||
|
if !msg.RecursionDesired {
|
||||||
|
t.Error("RecursionDesired should be true")
|
||||||
|
}
|
||||||
|
if opt := msg.IsEdns0(); opt == nil {
|
||||||
|
t.Error("expected EDNS0 OPT record")
|
||||||
|
} else if opt.UDPSize() != 2048 {
|
||||||
|
t.Errorf("EDNS0 UDPSize = %d, want 2048", opt.UDPSize())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildQueryFqdn(t *testing.T) {
|
||||||
|
msg := buildQuery("example.com", TypeA, 4096)
|
||||||
|
q := msg.Question[0]
|
||||||
|
if q.Name != "example.com." {
|
||||||
|
t.Errorf("Fqdn not applied: got %q, want %q", q.Name, "example.com.")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestQueryWithExchangeSuccess(t *testing.T) {
|
||||||
|
expectedResp := new(dns.Msg)
|
||||||
|
expectedResp.SetReply(new(dns.Msg))
|
||||||
|
expectedResp.Answer = append(expectedResp.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"),
|
||||||
|
})
|
||||||
|
|
||||||
|
exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||||
|
return expectedResp.Copy(), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg := &QueryConfig{
|
||||||
|
UDPSize: 2048,
|
||||||
|
Timeout: 5,
|
||||||
|
Retries: 1,
|
||||||
|
UseTCP: false,
|
||||||
|
}
|
||||||
|
|
||||||
|
server := net.ParseIP("8.8.8.8")
|
||||||
|
resp, err := QueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
if len(resp.Answer) != 1 {
|
||||||
|
t.Fatalf("expected 1 answer, got %d", len(resp.Answer))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestQueryTCPFallbackOnTruncation(t *testing.T) {
|
||||||
|
truncatedResp := new(dns.Msg)
|
||||||
|
truncatedResp.Truncated = true
|
||||||
|
truncatedResp.SetReply(new(dns.Msg))
|
||||||
|
|
||||||
|
fullResp := new(dns.Msg)
|
||||||
|
fullResp.SetReply(new(dns.Msg))
|
||||||
|
fullResp.Answer = append(fullResp.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"),
|
||||||
|
})
|
||||||
|
|
||||||
|
var mu sync.Mutex
|
||||||
|
calls := []bool{}
|
||||||
|
exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||||
|
mu.Lock()
|
||||||
|
defer mu.Unlock()
|
||||||
|
calls = append(calls, useTCP)
|
||||||
|
if !useTCP {
|
||||||
|
return truncatedResp.Copy(), nil
|
||||||
|
}
|
||||||
|
return fullResp.Copy(), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg := &QueryConfig{
|
||||||
|
UDPSize: 2048,
|
||||||
|
Timeout: 5,
|
||||||
|
Retries: 1,
|
||||||
|
UseTCP: false,
|
||||||
|
}
|
||||||
|
|
||||||
|
server := net.ParseIP("8.8.8.8")
|
||||||
|
resp, err := QueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
if len(calls) != 2 {
|
||||||
|
t.Fatalf("expected 2 exchange calls (UDP then TCP), got %d", len(calls))
|
||||||
|
}
|
||||||
|
if calls[0] != false {
|
||||||
|
t.Error("first call should be UDP")
|
||||||
|
}
|
||||||
|
if calls[1] != true {
|
||||||
|
t.Error("second call should be TCP")
|
||||||
|
}
|
||||||
|
if len(resp.Answer) != 1 {
|
||||||
|
t.Fatalf("expected 1 answer from TCP fallback, got %d", len(resp.Answer))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestQueryAlwaysTCP(t *testing.T) {
|
||||||
|
resp := new(dns.Msg)
|
||||||
|
resp.SetReply(new(dns.Msg))
|
||||||
|
|
||||||
|
var mu sync.Mutex
|
||||||
|
calls := []bool{}
|
||||||
|
exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||||
|
mu.Lock()
|
||||||
|
defer mu.Unlock()
|
||||||
|
calls = append(calls, useTCP)
|
||||||
|
return resp.Copy(), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg := &QueryConfig{
|
||||||
|
UDPSize: 2048,
|
||||||
|
Timeout: 5,
|
||||||
|
Retries: 1,
|
||||||
|
UseTCP: true,
|
||||||
|
}
|
||||||
|
|
||||||
|
server := net.ParseIP("8.8.8.8")
|
||||||
|
_, err := QueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
if len(calls) != 1 {
|
||||||
|
t.Fatalf("expected 1 exchange call, got %d", len(calls))
|
||||||
|
}
|
||||||
|
if !calls[0] {
|
||||||
|
t.Error("expected TCP call when UseTCP is true")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestQueryRetriesOnFailure(t *testing.T) {
|
||||||
|
var mu sync.Mutex
|
||||||
|
callCount := 0
|
||||||
|
exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||||
|
mu.Lock()
|
||||||
|
callCount++
|
||||||
|
mu.Unlock()
|
||||||
|
return nil, errors.New("connection refused")
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg := &QueryConfig{
|
||||||
|
UDPSize: 2048,
|
||||||
|
Timeout: 5,
|
||||||
|
Retries: 3,
|
||||||
|
UseTCP: false,
|
||||||
|
}
|
||||||
|
|
||||||
|
server := net.ParseIP("8.8.8.8")
|
||||||
|
_, err := QueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn)
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("expected error after retries exhausted")
|
||||||
|
}
|
||||||
|
if callCount != 3 {
|
||||||
|
t.Errorf("expected 3 calls (retries exhausted), got %d", callCount)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestQueryContextCancellation(t *testing.T) {
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
cancel()
|
||||||
|
|
||||||
|
exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||||
|
return nil, ctx.Err()
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg := &QueryConfig{
|
||||||
|
UDPSize: 2048,
|
||||||
|
Timeout: 5,
|
||||||
|
Retries: 1,
|
||||||
|
UseTCP: false,
|
||||||
|
}
|
||||||
|
|
||||||
|
server := net.ParseIP("8.8.8.8")
|
||||||
|
_, err := QueryWithExchange(ctx, server, "example.com", TypeA, cfg, exchangeFn)
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("expected error on cancelled context")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestQueryNilConfigUsesDefaults(t *testing.T) {
|
||||||
|
resp := new(dns.Msg)
|
||||||
|
resp.SetReply(new(dns.Msg))
|
||||||
|
|
||||||
|
exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||||
|
return resp.Copy(), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
server := net.ParseIP("8.8.8.8")
|
||||||
|
_, err := QueryWithExchange(context.Background(), server, "example.com", TypeA, nil, exchangeFn)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error with nil config: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestQueryZeroValuesUseDefaults(t *testing.T) {
|
||||||
|
resp := new(dns.Msg)
|
||||||
|
resp.SetReply(new(dns.Msg))
|
||||||
|
|
||||||
|
exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||||
|
return resp.Copy(), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg := &QueryConfig{
|
||||||
|
UDPSize: 0,
|
||||||
|
Timeout: 0,
|
||||||
|
Retries: 1,
|
||||||
|
UseTCP: false,
|
||||||
|
}
|
||||||
|
|
||||||
|
server := net.ParseIP("8.8.8.8")
|
||||||
|
_, err := QueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestQueryTCPFallbackFailsThenRetries(t *testing.T) {
|
||||||
|
truncatedResp := new(dns.Msg)
|
||||||
|
truncatedResp.Truncated = true
|
||||||
|
truncatedResp.SetReply(new(dns.Msg))
|
||||||
|
|
||||||
|
var mu sync.Mutex
|
||||||
|
callCount := 0
|
||||||
|
exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||||
|
mu.Lock()
|
||||||
|
callCount++
|
||||||
|
mu.Unlock()
|
||||||
|
if !useTCP {
|
||||||
|
return truncatedResp.Copy(), nil
|
||||||
|
}
|
||||||
|
return nil, errors.New("tcp failed")
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg := &QueryConfig{
|
||||||
|
UDPSize: 2048,
|
||||||
|
Timeout: 5,
|
||||||
|
Retries: 2,
|
||||||
|
UseTCP: false,
|
||||||
|
}
|
||||||
|
|
||||||
|
server := net.ParseIP("8.8.8.8")
|
||||||
|
_, err := QueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn)
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("expected error when TCP fallback always fails")
|
||||||
|
}
|
||||||
|
if callCount != 4 {
|
||||||
|
t.Errorf("expected 4 calls (2 retries x UDP+TCP), got %d", callCount)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,44 @@
|
|||||||
|
package dns
|
||||||
|
|
||||||
|
import "github.com/miekg/dns"
|
||||||
|
|
||||||
|
const (
|
||||||
|
TypeA uint16 = dns.TypeA
|
||||||
|
TypeAAAA uint16 = dns.TypeAAAA
|
||||||
|
TypeNS uint16 = dns.TypeNS
|
||||||
|
TypeCNAME uint16 = dns.TypeCNAME
|
||||||
|
TypeSOA uint16 = dns.TypeSOA
|
||||||
|
TypeMX uint16 = dns.TypeMX
|
||||||
|
TypeTXT uint16 = dns.TypeTXT
|
||||||
|
TypeSRV uint16 = dns.TypeSRV
|
||||||
|
TypePTR uint16 = dns.TypePTR
|
||||||
|
TypeANY uint16 = dns.TypeANY
|
||||||
|
)
|
||||||
|
|
||||||
|
var QNameTypes = map[uint16]string{
|
||||||
|
TypeA: "A",
|
||||||
|
TypeAAAA: "AAAA",
|
||||||
|
TypeNS: "NS",
|
||||||
|
TypeCNAME: "CNAME",
|
||||||
|
TypeSOA: "SOA",
|
||||||
|
TypeMX: "MX",
|
||||||
|
TypeTXT: "TXT",
|
||||||
|
TypeSRV: "SRV",
|
||||||
|
TypePTR: "PTR",
|
||||||
|
TypeANY: "ANY",
|
||||||
|
}
|
||||||
|
|
||||||
|
func QNameType(qtype uint16) string {
|
||||||
|
if name, ok := QNameTypes[qtype]; ok {
|
||||||
|
return name
|
||||||
|
}
|
||||||
|
return dns.TypeToString[qtype]
|
||||||
|
}
|
||||||
|
|
||||||
|
func DefaultEDNS0UDPSize() int {
|
||||||
|
return 2048
|
||||||
|
}
|
||||||
|
|
||||||
|
func MinEDNS0UDPSize() int {
|
||||||
|
return 512
|
||||||
|
}
|
||||||
@@ -0,0 +1,76 @@
|
|||||||
|
package dns
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestQNameType(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
qtype uint16
|
||||||
|
expect string
|
||||||
|
}{
|
||||||
|
{"A record", TypeA, "A"},
|
||||||
|
{"AAAA record", TypeAAAA, "AAAA"},
|
||||||
|
{"NS record", TypeNS, "NS"},
|
||||||
|
{"CNAME record", TypeCNAME, "CNAME"},
|
||||||
|
{"SOA record", TypeSOA, "SOA"},
|
||||||
|
{"MX record", TypeMX, "MX"},
|
||||||
|
{"TXT record", TypeTXT, "TXT"},
|
||||||
|
{"SRV record", TypeSRV, "SRV"},
|
||||||
|
{"PTR record", TypePTR, "PTR"},
|
||||||
|
{"ANY record", TypeANY, "ANY"},
|
||||||
|
{"unknown type", uint16(9999), ""},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
got := QNameType(tt.qtype)
|
||||||
|
if got != tt.expect {
|
||||||
|
t.Errorf("QNameType(%d) = %q, want %q", tt.qtype, got, tt.expect)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConstantsMatchMiekg(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
local uint16
|
||||||
|
}{
|
||||||
|
{"TypeA", TypeA},
|
||||||
|
{"TypeAAAA", TypeAAAA},
|
||||||
|
{"TypeNS", TypeNS},
|
||||||
|
{"TypeCNAME", TypeCNAME},
|
||||||
|
{"TypeSOA", TypeSOA},
|
||||||
|
{"TypeMX", TypeMX},
|
||||||
|
{"TypeTXT", TypeTXT},
|
||||||
|
{"TypeSRV", TypeSRV},
|
||||||
|
{"TypePTR", TypePTR},
|
||||||
|
{"TypeANY", TypeANY},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
mapped, ok := QNameTypes[tt.local]
|
||||||
|
if !ok {
|
||||||
|
t.Errorf("QNameTypes missing entry for %s (%d)", tt.name, tt.local)
|
||||||
|
}
|
||||||
|
if mapped != QNameType(tt.local) {
|
||||||
|
t.Errorf("QNameType(%d) = %q, QNameTypes[%d] = %q", tt.local, QNameType(tt.local), tt.local, mapped)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDefaultEDNS0UDPSize(t *testing.T) {
|
||||||
|
if got := DefaultEDNS0UDPSize(); got != 2048 {
|
||||||
|
t.Errorf("DefaultEDNS0UDPSize() = %d, want 2048", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMinEDNS0UDPSize(t *testing.T) {
|
||||||
|
if got := MinEDNS0UDPSize(); got != 512 {
|
||||||
|
t.Errorf("MinEDNS0UDPSize() = %d, want 512", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user