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

Reviewed-on: http://gitea.hansenits.com.au/hits/ExploreDNS/pulls/3
This commit was merged in pull request #3.
This commit is contained in:
2026-06-05 05:37:38 +00:00
8 changed files with 1127 additions and 1 deletions
+2 -1
View File
@@ -2,8 +2,9 @@ module github.com/hits/ExploreDNS
go 1.25.6
require github.com/miekg/dns v1.1.72
require (
github.com/miekg/dns v1.1.72 // indirect
golang.org/x/mod v0.31.0 // indirect
golang.org/x/net v0.48.0 // indirect
golang.org/x/sync v0.19.0 // indirect
+2
View File
@@ -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/go.mod h1:+EuEPhdHOsfk6Wk5TT2CzssZdqkmFhf8r+aVyDEToIs=
golang.org/x/mod v0.31.0 h1:HaW9xtz0+kOcWKwli0ZXy79Ix+UW/vOfmWI5QVd2tgI=
+199
View File
@@ -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(),
)
}
+381
View File
@@ -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()
}
+134
View File
@@ -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
}
+289
View File
@@ -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)
}
}
+44
View File
@@ -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
}
+76
View File
@@ -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)
}
}