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") } }