From fce79206110d1fa99072dc07381553260fe9763d Mon Sep 17 00:00:00 2001 From: Gary Date: Fri, 5 Jun 2026 15:14:56 +1000 Subject: [PATCH] feat: add DNS query/response layer with miekg/dns 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 --- go.mod | 13 ++ go.sum | 14 ++ internal/dns/decode.go | 199 +++++++++++++++++++ internal/dns/decode_test.go | 381 ++++++++++++++++++++++++++++++++++++ internal/dns/query.go | 134 +++++++++++++ internal/dns/query_test.go | 289 +++++++++++++++++++++++++++ internal/dns/types.go | 44 +++++ internal/dns/types_test.go | 76 +++++++ 8 files changed, 1150 insertions(+) create mode 100644 go.mod create mode 100644 go.sum create mode 100644 internal/dns/decode.go create mode 100644 internal/dns/decode_test.go create mode 100644 internal/dns/query.go create mode 100644 internal/dns/query_test.go create mode 100644 internal/dns/types.go create mode 100644 internal/dns/types_test.go diff --git a/go.mod b/go.mod new file mode 100644 index 0000000..29ccbfa --- /dev/null +++ b/go.mod @@ -0,0 +1,13 @@ +module github.com/hits/ExploreDNS + +go 1.25.6 + +require github.com/miekg/dns v1.1.72 + +require ( + golang.org/x/mod v0.31.0 // indirect + golang.org/x/net v0.48.0 // indirect + golang.org/x/sync v0.19.0 // indirect + golang.org/x/sys v0.39.0 // indirect + golang.org/x/tools v0.40.0 // indirect +) diff --git a/go.sum b/go.sum new file mode 100644 index 0000000..92364b3 --- /dev/null +++ b/go.sum @@ -0,0 +1,14 @@ +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= +golang.org/x/mod v0.31.0/go.mod h1:43JraMp9cGx1Rx3AqioxrbrhNsLl2l/iNAvuBkrezpg= +golang.org/x/net v0.48.0 h1:zyQRTTrjc33Lhh0fBgT/H3oZq9WuvRR5gPC70xpDiQU= +golang.org/x/net v0.48.0/go.mod h1:+ndRgGjkh8FGtu1w1FGbEC31if4VrNVMuKTgcAAnQRY= +golang.org/x/sync v0.19.0 h1:vV+1eWNmZ5geRlYjzm2adRgW2/mcpevXNg50YZtPCE4= +golang.org/x/sync v0.19.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI= +golang.org/x/sys v0.39.0 h1:CvCKL8MeisomCi6qNZ+wbb0DN9E5AATixKsvNtMoMFk= +golang.org/x/sys v0.39.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks= +golang.org/x/tools v0.40.0 h1:yLkxfA+Qnul4cs9QA3KnlFu0lVmd8JJfoq+E41uSutA= +golang.org/x/tools v0.40.0/go.mod h1:Ik/tzLRlbscWpqqMRjyWYDisX8bG13FrdXp3o4Sr9lc= diff --git a/internal/dns/decode.go b/internal/dns/decode.go new file mode 100644 index 0000000..a909f24 --- /dev/null +++ b/internal/dns/decode.go @@ -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(), + ) +} diff --git a/internal/dns/decode_test.go b/internal/dns/decode_test.go new file mode 100644 index 0000000..121f817 --- /dev/null +++ b/internal/dns/decode_test.go @@ -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() +} diff --git a/internal/dns/query.go b/internal/dns/query.go new file mode 100644 index 0000000..e219f8a --- /dev/null +++ b/internal/dns/query.go @@ -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 +} diff --git a/internal/dns/query_test.go b/internal/dns/query_test.go new file mode 100644 index 0000000..5c18f27 --- /dev/null +++ b/internal/dns/query_test.go @@ -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) + } +} diff --git a/internal/dns/types.go b/internal/dns/types.go new file mode 100644 index 0000000..f17064a --- /dev/null +++ b/internal/dns/types.go @@ -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 +} diff --git a/internal/dns/types_test.go b/internal/dns/types_test.go new file mode 100644 index 0000000..bd75ab7 --- /dev/null +++ b/internal/dns/types_test.go @@ -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) + } +}