305 lines
8.5 KiB
Go
305 lines
8.5 KiB
Go
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, systemResolver(), "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")
|
|
}
|
|
}
|