test: merge PR #13 test coverage improvements into main
Resolves merge conflicts between Phase 4.2 comprehensive test suite and the test coverage improvement branch: - config_test.go: take PR's better table-driven tests + keep main's extra tests - coverage_test.go: keep main's Phase 4.2 comprehensive tests Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Co-authored-by: multica-agent <github@multica.ai>
This commit is contained in:
@@ -0,0 +1,242 @@
|
||||
package dns
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"github.com/miekg/dns"
|
||||
)
|
||||
|
||||
func TestIterativeQueryWithExchangeSuccess(t *testing.T) {
|
||||
resp := new(dns.Msg)
|
||||
resp.SetReply(new(dns.Msg))
|
||||
resp.RecursionDesired = false
|
||||
resp.Answer = append(resp.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) {
|
||||
if msg.RecursionDesired {
|
||||
t.Error("iterative query should have RecursionDesired=false")
|
||||
}
|
||||
return resp.Copy(), nil
|
||||
}
|
||||
|
||||
cfg := &QueryConfig{
|
||||
UDPSize: 2048,
|
||||
Retries: 1,
|
||||
UseTCP: false,
|
||||
}
|
||||
|
||||
server := net.ParseIP("8.8.8.8")
|
||||
result, err := IterativeQueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if len(result.Answer) != 1 {
|
||||
t.Fatalf("expected 1 answer, got %d", len(result.Answer))
|
||||
}
|
||||
}
|
||||
|
||||
func TestIterativeQueryWithExchangeNilConfig(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 := IterativeQueryWithExchange(context.Background(), server, "example.com", TypeA, nil, exchangeFn)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error with nil config: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestIterativeQueryWithExchangeNilResponse(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, nil
|
||||
}
|
||||
|
||||
cfg := &QueryConfig{
|
||||
UDPSize: 2048,
|
||||
Retries: 2,
|
||||
UseTCP: false,
|
||||
}
|
||||
|
||||
server := net.ParseIP("8.8.8.8")
|
||||
_, err := IterativeQueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn)
|
||||
if err == nil {
|
||||
t.Fatal("expected error when response is nil")
|
||||
}
|
||||
if callCount != 2 {
|
||||
t.Errorf("expected 2 attempts, got %d", callCount)
|
||||
}
|
||||
}
|
||||
|
||||
func TestIterativeQueryWithExchangeTCPFallback(t *testing.T) {
|
||||
truncated := new(dns.Msg)
|
||||
truncated.Truncated = true
|
||||
truncated.SetReply(new(dns.Msg))
|
||||
|
||||
full := new(dns.Msg)
|
||||
full.SetReply(new(dns.Msg))
|
||||
full.Answer = append(full.Answer, &dns.A{
|
||||
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA},
|
||||
A: net.ParseIP("1.2.3.4"),
|
||||
})
|
||||
|
||||
var mu sync.Mutex
|
||||
calls := []bool{}
|
||||
exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||
mu.Lock()
|
||||
calls = append(calls, useTCP)
|
||||
mu.Unlock()
|
||||
if !useTCP {
|
||||
return truncated.Copy(), nil
|
||||
}
|
||||
return full.Copy(), nil
|
||||
}
|
||||
|
||||
cfg := &QueryConfig{
|
||||
UDPSize: 2048,
|
||||
Retries: 1,
|
||||
UseTCP: false,
|
||||
AllowTCP: true,
|
||||
}
|
||||
|
||||
server := net.ParseIP("8.8.8.8")
|
||||
result, err := IterativeQueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if len(calls) != 2 {
|
||||
t.Errorf("expected 2 calls (UDP + TCP), got %d", len(calls))
|
||||
}
|
||||
if len(result.Answer) != 1 {
|
||||
t.Fatalf("expected 1 answer, got %d", len(result.Answer))
|
||||
}
|
||||
}
|
||||
|
||||
func TestIterativeQueryWithExchangeAlwaysTCP(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()
|
||||
calls = append(calls, useTCP)
|
||||
mu.Unlock()
|
||||
return resp.Copy(), nil
|
||||
}
|
||||
|
||||
cfg := &QueryConfig{
|
||||
UDPSize: 2048,
|
||||
Retries: 1,
|
||||
UseTCP: true,
|
||||
}
|
||||
|
||||
server := net.ParseIP("8.8.8.8")
|
||||
_, err := IterativeQueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if len(calls) != 1 || !calls[0] {
|
||||
t.Errorf("expected single TCP call, got %v", calls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestIterativeQueryWithExchangeRetriesOnFailure(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,
|
||||
Retries: 3,
|
||||
UseTCP: false,
|
||||
}
|
||||
|
||||
server := net.ParseIP("8.8.8.8")
|
||||
_, err := IterativeQueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn)
|
||||
if err == nil {
|
||||
t.Fatal("expected error")
|
||||
}
|
||||
if callCount != 3 {
|
||||
t.Errorf("expected 3 calls, got %d", callCount)
|
||||
}
|
||||
}
|
||||
|
||||
func TestIterativeQueryWithExchangeContextCancel(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,
|
||||
Retries: 1,
|
||||
UseTCP: false,
|
||||
}
|
||||
|
||||
server := net.ParseIP("8.8.8.8")
|
||||
_, err := IterativeQueryWithExchange(ctx, server, "example.com", TypeA, cfg, exchangeFn)
|
||||
if err == nil {
|
||||
t.Fatal("expected error on cancelled context")
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractNSNames(t *testing.T) {
|
||||
rrs := []dns.RR{
|
||||
&dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS}, Ns: "a.root-servers.net."},
|
||||
&dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS}, Ns: "b.root-servers.net."},
|
||||
}
|
||||
names := extractNSNames(rrs)
|
||||
if len(names) != 2 {
|
||||
t.Fatalf("expected 2 names, got %d", len(names))
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractNSNamesEmpty(t *testing.T) {
|
||||
names := extractNSNames(nil)
|
||||
if len(names) != 0 {
|
||||
t.Errorf("expected 0 names, got %d", len(names))
|
||||
}
|
||||
}
|
||||
|
||||
func TestIterativeQueryWithExchangeZeroUDPSize(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, // should use default
|
||||
Retries: 1,
|
||||
}
|
||||
|
||||
server := net.ParseIP("8.8.8.8")
|
||||
_, err := IterativeQueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,304 @@
|
||||
package dns
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/miekg/dns"
|
||||
)
|
||||
|
||||
// startTestDNSServer starts a local DNS server on a random port and returns the address and a stop function.
|
||||
func startTestDNSServer(t *testing.T, handler dns.HandlerFunc) (string, func()) {
|
||||
t.Helper()
|
||||
|
||||
pc, err := net.ListenPacket("udp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Skipf("cannot start test DNS server: %v", err)
|
||||
}
|
||||
addr := pc.LocalAddr().String()
|
||||
|
||||
mux := dns.NewServeMux()
|
||||
mux.HandleFunc(".", handler)
|
||||
|
||||
srv := &dns.Server{
|
||||
PacketConn: pc,
|
||||
Net: "udp",
|
||||
Handler: mux,
|
||||
}
|
||||
|
||||
started := make(chan struct{})
|
||||
srv.NotifyStartedFunc = func() { close(started) }
|
||||
|
||||
go func() {
|
||||
_ = srv.ActivateAndServe()
|
||||
}()
|
||||
|
||||
select {
|
||||
case <-started:
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Skip("test DNS server did not start in time")
|
||||
}
|
||||
|
||||
return addr, func() { _ = srv.Shutdown() }
|
||||
}
|
||||
|
||||
func TestQueryUsesRealExchange(t *testing.T) {
|
||||
addr, stop := startTestDNSServer(t, func(w dns.ResponseWriter, r *dns.Msg) {
|
||||
resp := new(dns.Msg)
|
||||
resp.SetReply(r)
|
||||
resp.Answer = append(resp.Answer, &dns.A{
|
||||
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300},
|
||||
A: net.ParseIP("1.2.3.4"),
|
||||
})
|
||||
_ = w.WriteMsg(resp)
|
||||
})
|
||||
defer stop()
|
||||
|
||||
host, portStr, err := net.SplitHostPort(addr)
|
||||
if err != nil {
|
||||
t.Fatalf("parse addr: %v", err)
|
||||
}
|
||||
var port int
|
||||
fmt.Sscanf(portStr, "%d", &port)
|
||||
|
||||
// Patch the realExchange to use the test server by using QueryWithExchange with a custom exchangeFn.
|
||||
// Since we can't inject into Query directly, use realExchangeWithPort for test.
|
||||
serverIP := net.ParseIP(host)
|
||||
cfg := DefaultQueryConfig()
|
||||
cfg.Retries = 1
|
||||
|
||||
// Test QueryWithExchange (already covered), but now test Query+realExchange flow via
|
||||
// a patched exchange that routes to our test server port.
|
||||
patchedExchange := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||
c := &dns.Client{Net: "udp", ReadTimeout: 3 * time.Second, WriteTimeout: 3 * time.Second}
|
||||
r, _, err := c.ExchangeContext(ctx, msg, fmt.Sprintf("%s:%d", host, port))
|
||||
return r, err
|
||||
}
|
||||
|
||||
resp, err := QueryWithExchange(context.Background(), serverIP, "example.com", TypeA, cfg, patchedExchange)
|
||||
if err != nil {
|
||||
t.Fatalf("QueryWithExchange: %v", err)
|
||||
}
|
||||
if len(resp.Answer) == 0 {
|
||||
t.Fatal("expected at least 1 answer")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRealExchangeViaDirect(t *testing.T) {
|
||||
// Test realExchange directly via the exported Query function
|
||||
// by using a server that will respond or fail quickly.
|
||||
// We use a loopback address with a timeout to exercise code paths.
|
||||
addr, stop := startTestDNSServer(t, func(w dns.ResponseWriter, r *dns.Msg) {
|
||||
resp := new(dns.Msg)
|
||||
resp.SetReply(r)
|
||||
resp.Answer = append(resp.Answer, &dns.A{
|
||||
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300},
|
||||
A: net.ParseIP("5.6.7.8"),
|
||||
})
|
||||
_ = w.WriteMsg(resp)
|
||||
})
|
||||
defer stop()
|
||||
|
||||
host, portStr, _ := net.SplitHostPort(addr)
|
||||
serverIP := net.ParseIP(host)
|
||||
|
||||
// Exercise realExchange via Query — we need a way to target the test port.
|
||||
// Use a custom exchange that calls through realExchange-like logic.
|
||||
cfg := DefaultQueryConfig()
|
||||
cfg.Retries = 1
|
||||
|
||||
resp, err := QueryWithExchange(context.Background(), serverIP, "example.com", TypeA, cfg,
|
||||
func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||
targetAddr := fmt.Sprintf("%s:%s", host, portStr)
|
||||
c := &dns.Client{Net: "udp", ReadTimeout: 3 * time.Second, WriteTimeout: 3 * time.Second}
|
||||
r, _, e := c.ExchangeContext(ctx, msg, targetAddr)
|
||||
return r, e
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if len(resp.Answer) == 0 {
|
||||
t.Fatal("expected answers")
|
||||
}
|
||||
}
|
||||
|
||||
func TestQueryFunctionDirectly(t *testing.T) {
|
||||
// Exercise Query() itself (which calls realExchange) by using 127.0.0.1:53.
|
||||
// The test skips if no local DNS is available.
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
|
||||
defer cancel()
|
||||
|
||||
server := net.ParseIP("127.0.0.1")
|
||||
cfg := DefaultQueryConfig()
|
||||
cfg.Retries = 1
|
||||
cfg.Timeout = 2 * time.Second
|
||||
|
||||
_, err := Query(ctx, server, ".", TypeNS, cfg)
|
||||
if err != nil {
|
||||
t.Skipf("skipping (no local DNS at 127.0.0.1:53): %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestIterativeQueryDirectly(t *testing.T) {
|
||||
// Exercise IterativeQuery() itself (which calls realExchange) by using 127.0.0.1:53.
|
||||
// The test skips if no local DNS is available.
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
|
||||
defer cancel()
|
||||
|
||||
server := net.ParseIP("127.0.0.1")
|
||||
cfg := DefaultQueryConfig()
|
||||
cfg.Retries = 1
|
||||
cfg.Timeout = 2 * time.Second
|
||||
|
||||
_, err := IterativeQuery(ctx, server, ".", TypeNS, cfg)
|
||||
if err != nil {
|
||||
t.Skipf("skipping (no local DNS at 127.0.0.1:53): %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBasicResolverQuery(t *testing.T) {
|
||||
// Exercise BasicResolver.Query() which calls Query() → realExchange.
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
|
||||
defer cancel()
|
||||
|
||||
br := NewBasicResolver()
|
||||
server := net.ParseIP("127.0.0.1")
|
||||
cfg := DefaultQueryConfig()
|
||||
cfg.Retries = 1
|
||||
cfg.Timeout = 2 * time.Second
|
||||
|
||||
_, err := br.Query(ctx, server, ".", TypeNS, cfg)
|
||||
if err != nil {
|
||||
t.Skipf("skipping (no local DNS at 127.0.0.1:53): %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiscoverAllRoots(t *testing.T) {
|
||||
// discoverAllRoots calls queryResolver(ctx, "127.0.0.1:53", ...)
|
||||
// Skip if local DNS is not available.
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
|
||||
cfg := &RootDiscoveryConfig{
|
||||
AllRoots: true,
|
||||
IncludeAAAA: false,
|
||||
}
|
||||
servers, err := DiscoverRoots(ctx, cfg)
|
||||
if err != nil {
|
||||
t.Skipf("skipping (no local DNS available): %v", err)
|
||||
}
|
||||
if len(servers) == 0 {
|
||||
t.Fatal("expected at least one root server from discoverAllRoots")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiscoverAllRootsWithAAAA(t *testing.T) {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
|
||||
cfg := &RootDiscoveryConfig{
|
||||
AllRoots: true,
|
||||
IncludeAAAA: true,
|
||||
}
|
||||
servers, err := DiscoverRoots(ctx, cfg)
|
||||
if err != nil {
|
||||
t.Skipf("skipping (no local DNS available): %v", err)
|
||||
}
|
||||
if len(servers) == 0 {
|
||||
t.Fatal("expected root servers with AAAA")
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveRootServerDirect(t *testing.T) {
|
||||
// Calls resolveRootServer directly (unexported, but in same package).
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
|
||||
servers, err := resolveRootServer(ctx, "a.root-servers.net.", false)
|
||||
if err != nil {
|
||||
t.Skipf("skipping (no local DNS): %v", err)
|
||||
}
|
||||
if len(servers) == 0 || len(servers[0].IPv4) == 0 {
|
||||
t.Fatal("expected IPv4 address for a.root-servers.net.")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiscoverSingleRootWithAAAA(t *testing.T) {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
|
||||
cfg := &RootDiscoveryConfig{
|
||||
AllRoots: false,
|
||||
IncludeAAAA: true,
|
||||
}
|
||||
servers, err := DiscoverRoots(ctx, cfg)
|
||||
if err != nil {
|
||||
t.Skipf("skipping (no local DNS available): %v", err)
|
||||
}
|
||||
if len(servers) == 0 {
|
||||
t.Fatal("expected at least one root server")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRealExchangeTCPPath(t *testing.T) {
|
||||
// Test the TCP path of realExchange via a test server
|
||||
tcpAddr := ""
|
||||
listener, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Skipf("cannot start TCP test server: %v", err)
|
||||
}
|
||||
tcpAddr = listener.Addr().String()
|
||||
|
||||
mux := dns.NewServeMux()
|
||||
mux.HandleFunc(".", func(w dns.ResponseWriter, r *dns.Msg) {
|
||||
resp := new(dns.Msg)
|
||||
resp.SetReply(r)
|
||||
resp.Answer = append(resp.Answer, &dns.A{
|
||||
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300},
|
||||
A: net.ParseIP("9.9.9.9"),
|
||||
})
|
||||
_ = w.WriteMsg(resp)
|
||||
})
|
||||
|
||||
srv := &dns.Server{
|
||||
Listener: listener,
|
||||
Net: "tcp",
|
||||
Handler: mux,
|
||||
}
|
||||
|
||||
started := make(chan struct{})
|
||||
srv.NotifyStartedFunc = func() { close(started) }
|
||||
|
||||
go func() { _ = srv.ActivateAndServe() }()
|
||||
|
||||
select {
|
||||
case <-started:
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Skip("TCP DNS server didn't start")
|
||||
}
|
||||
defer srv.Shutdown()
|
||||
|
||||
host, portStr, _ := net.SplitHostPort(tcpAddr)
|
||||
serverIP := net.ParseIP(host)
|
||||
|
||||
cfg := DefaultQueryConfig()
|
||||
cfg.UseTCP = true
|
||||
cfg.Retries = 1
|
||||
|
||||
resp, err := QueryWithExchange(context.Background(), serverIP, "example.com", TypeA, cfg,
|
||||
func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||
targetAddr := fmt.Sprintf("%s:%s", host, portStr)
|
||||
c := &dns.Client{Net: "tcp", ReadTimeout: 3 * time.Second, WriteTimeout: 3 * time.Second}
|
||||
r, _, e := c.ExchangeContext(ctx, msg, targetAddr)
|
||||
return r, e
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("TCP query: %v", err)
|
||||
}
|
||||
if len(resp.Answer) == 0 {
|
||||
t.Fatal("expected answers")
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user