TestClientAgainstLocalServer captured RecursionDesired into a plain bool from the miekg server handler goroutine and read it from the test goroutine; the UDP round-trip gives no happens-before edge, so CI's -race run flagged it. Use atomic.Bool. Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
137 lines
3.6 KiB
Go
137 lines
3.6 KiB
Go
package dns
|
|
|
|
import (
|
|
"context"
|
|
"net"
|
|
"sync/atomic"
|
|
"testing"
|
|
"time"
|
|
|
|
"github.com/miekg/dns"
|
|
)
|
|
|
|
// startTestDNSServer starts a loopback DNS server on a random port and
|
|
// returns its address. No test in this file touches the real network.
|
|
func startTestDNSServer(t *testing.T, network string, handler dns.HandlerFunc) string {
|
|
t.Helper()
|
|
|
|
mux := dns.NewServeMux()
|
|
mux.HandleFunc(".", handler)
|
|
|
|
srv := &dns.Server{Net: network, Handler: mux}
|
|
var addr string
|
|
|
|
switch network {
|
|
case "udp":
|
|
pc, err := net.ListenPacket("udp", "127.0.0.1:0")
|
|
if err != nil {
|
|
t.Skipf("cannot start test DNS server: %v", err)
|
|
}
|
|
srv.PacketConn = pc
|
|
addr = pc.LocalAddr().String()
|
|
case "tcp":
|
|
l, err := net.Listen("tcp", "127.0.0.1:0")
|
|
if err != nil {
|
|
t.Skipf("cannot start test DNS server: %v", err)
|
|
}
|
|
srv.Listener = l
|
|
addr = l.Addr().String()
|
|
}
|
|
|
|
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")
|
|
}
|
|
|
|
t.Cleanup(func() { _ = srv.Shutdown() })
|
|
return addr
|
|
}
|
|
|
|
func aHandler(ip string) dns.HandlerFunc {
|
|
return 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: r.Question[0].Name, Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300},
|
|
A: net.ParseIP(ip),
|
|
})
|
|
_ = w.WriteMsg(resp)
|
|
}
|
|
}
|
|
|
|
func TestRealExchangeUDPHostPort(t *testing.T) {
|
|
addr := startTestDNSServer(t, "udp", aHandler("1.2.3.4"))
|
|
|
|
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
|
|
defer cancel()
|
|
|
|
// realExchange must honour an explicit host:port (used by root discovery
|
|
// upstream resolvers).
|
|
resp, err := realExchange(ctx, addr, buildQuery("example.com.", TypeA, 2048), false)
|
|
if err != nil {
|
|
t.Fatalf("realExchange: %v", err)
|
|
}
|
|
if len(resp.Answer) == 0 {
|
|
t.Fatal("expected answers")
|
|
}
|
|
}
|
|
|
|
func TestRealExchangeTCP(t *testing.T) {
|
|
addr := startTestDNSServer(t, "tcp", aHandler("9.9.9.9"))
|
|
|
|
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
|
|
defer cancel()
|
|
|
|
resp, err := realExchange(ctx, addr, buildQuery("example.com.", TypeA, 2048), true)
|
|
if err != nil {
|
|
t.Fatalf("realExchange TCP: %v", err)
|
|
}
|
|
if len(resp.Answer) == 0 {
|
|
t.Fatal("expected answers")
|
|
}
|
|
}
|
|
|
|
func TestRealExchangeUnreachable(t *testing.T) {
|
|
ctx, cancel := context.WithTimeout(context.Background(), 200*time.Millisecond)
|
|
defer cancel()
|
|
|
|
_, err := realExchange(ctx, "127.0.0.1:1", buildQuery("example.com.", TypeA, 2048), false)
|
|
if err == nil {
|
|
t.Fatal("expected error for unreachable server")
|
|
}
|
|
}
|
|
|
|
func TestClientAgainstLocalServer(t *testing.T) {
|
|
// Written by the server handler goroutine, read by the test goroutine;
|
|
// the UDP round-trip provides no happens-before edge, so use an atomic.
|
|
var sawRD atomic.Bool
|
|
addr := startTestDNSServer(t, "udp", func(w dns.ResponseWriter, r *dns.Msg) {
|
|
sawRD.Store(r.RecursionDesired)
|
|
aHandler("5.6.7.8")(w, r)
|
|
})
|
|
|
|
// Route the client's exchange to the test server's port while still
|
|
// exercising realExchange.
|
|
c := NewClient(&QueryConfig{Retries: 1, Timeout: 2 * time.Second},
|
|
func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
|
return realExchange(ctx, addr, msg, useTCP)
|
|
})
|
|
|
|
resp, _, err := c.Query(context.Background(), net.ParseIP("127.0.0.1"), "example.com", TypeA)
|
|
if err != nil {
|
|
t.Fatalf("Client.Query: %v", err)
|
|
}
|
|
if len(resp.Answer) == 0 {
|
|
t.Fatal("expected answers")
|
|
}
|
|
if sawRD.Load() {
|
|
t.Error("wire query must have RD=0")
|
|
}
|
|
}
|