471 lines
12 KiB
Go
471 lines
12 KiB
Go
package dns
|
|
|
|
import (
|
|
"context"
|
|
"net"
|
|
"testing"
|
|
"time"
|
|
|
|
"github.com/miekg/dns"
|
|
)
|
|
|
|
func TestRootServerAllIPs(t *testing.T) {
|
|
rs := RootServer{
|
|
Name: "a.root-servers.net.",
|
|
IPv4: []net.IP{net.ParseIP("198.41.0.4")},
|
|
IPv6: []net.IP{net.ParseIP("2001:503:ba3e::2:30")},
|
|
}
|
|
|
|
t.Run("IPv4 only", func(t *testing.T) {
|
|
ips := rs.AllIPs(false)
|
|
if len(ips) != 1 {
|
|
t.Fatalf("expected 1 IP, got %d", len(ips))
|
|
}
|
|
if !ips[0].Equal(net.ParseIP("198.41.0.4")) {
|
|
t.Errorf("got %v, want 198.41.0.4", ips[0])
|
|
}
|
|
})
|
|
|
|
t.Run("IPv4 and IPv6", func(t *testing.T) {
|
|
ips := rs.AllIPs(true)
|
|
if len(ips) != 2 {
|
|
t.Fatalf("expected 2 IPs, got %d", len(ips))
|
|
}
|
|
})
|
|
|
|
t.Run("no addresses", func(t *testing.T) {
|
|
empty := RootServer{Name: "empty.root-servers.net."}
|
|
ips := empty.AllIPs(false)
|
|
if len(ips) != 0 {
|
|
t.Errorf("expected 0 IPs, got %d", len(ips))
|
|
}
|
|
})
|
|
}
|
|
|
|
func TestDefaultRootDiscoveryConfig(t *testing.T) {
|
|
cfg := DefaultRootDiscoveryConfig()
|
|
if cfg.AllRoots {
|
|
t.Error("AllRoots should be false by default")
|
|
}
|
|
if cfg.IncludeAAAA {
|
|
t.Error("IncludeAAAA should be false by default")
|
|
}
|
|
if cfg.Server != "" {
|
|
t.Error("Server should be empty by default")
|
|
}
|
|
}
|
|
|
|
func TestNormalizeServerName(t *testing.T) {
|
|
tests := []struct {
|
|
input string
|
|
want string
|
|
}{
|
|
{"a.root-servers.net.", "a.root-servers.net"},
|
|
{"A.ROOT-SERVERS.NET.", "a.root-servers.net"},
|
|
{"b.root-servers.net", "b.root-servers.net"},
|
|
{"root-servers.net.", "root-servers.net"},
|
|
}
|
|
|
|
for _, tt := range tests {
|
|
t.Run(tt.input, func(t *testing.T) {
|
|
got := normalizeServerName(tt.input)
|
|
if got != tt.want {
|
|
t.Errorf("normalizeServerName(%q) = %q, want %q", tt.input, got, tt.want)
|
|
}
|
|
})
|
|
}
|
|
}
|
|
|
|
func TestExtractNSRecords(t *testing.T) {
|
|
t.Run("empty", func(t *testing.T) {
|
|
names := extractNSRecords(nil)
|
|
if len(names) != 0 {
|
|
t.Errorf("expected 0 names, got %d", len(names))
|
|
}
|
|
})
|
|
|
|
t.Run("NS records", func(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 := extractNSRecords(rrs)
|
|
if len(names) != 2 {
|
|
t.Fatalf("expected 2 names, got %d", len(names))
|
|
}
|
|
if names[0] != "a.root-servers.net." {
|
|
t.Errorf("got %q, want a.root-servers.net.", names[0])
|
|
}
|
|
})
|
|
|
|
t.Run("dedup", func(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: "a.root-servers.net."},
|
|
}
|
|
names := extractNSRecords(rrs)
|
|
if len(names) != 1 {
|
|
t.Errorf("expected 1 deduped name, got %d", len(names))
|
|
}
|
|
|
|
})
|
|
}
|
|
|
|
func TestExtractIPsFromAnswer(t *testing.T) {
|
|
t.Run("empty", func(t *testing.T) {
|
|
ips := extractIPsFromAnswer(nil, dns.TypeA)
|
|
if len(ips) != 0 {
|
|
t.Errorf("expected 0 IPs, got %d", len(ips))
|
|
}
|
|
})
|
|
|
|
t.Run("A records", func(t *testing.T) {
|
|
rrs := []dns.RR{
|
|
&dns.A{Hdr: dns.RR_Header{Rrtype: dns.TypeA}, A: net.ParseIP("198.41.0.4")},
|
|
&dns.A{Hdr: dns.RR_Header{Rrtype: dns.TypeA}, A: net.ParseIP("199.9.14.201")},
|
|
}
|
|
ips := extractIPsFromAnswer(rrs, dns.TypeA)
|
|
if len(ips) != 2 {
|
|
t.Fatalf("expected 2 IPs, got %d", len(ips))
|
|
}
|
|
})
|
|
|
|
t.Run("AAAA records", func(t *testing.T) {
|
|
rrs := []dns.RR{
|
|
&dns.AAAA{Hdr: dns.RR_Header{Rrtype: dns.TypeAAAA}, AAAA: net.ParseIP("2001:503:ba3e::2:30")},
|
|
}
|
|
ips := extractIPsFromAnswer(rrs, dns.TypeAAAA)
|
|
if len(ips) != 1 {
|
|
t.Fatalf("expected 1 IP, got %d", len(ips))
|
|
}
|
|
})
|
|
|
|
t.Run("filter by type", func(t *testing.T) {
|
|
rrs := []dns.RR{
|
|
&dns.A{Hdr: dns.RR_Header{Rrtype: dns.TypeA}, A: net.ParseIP("198.41.0.4")},
|
|
&dns.AAAA{Hdr: dns.RR_Header{Rrtype: dns.TypeAAAA}, AAAA: net.ParseIP("2001:503:ba3e::2:30")},
|
|
}
|
|
ips := extractIPsFromAnswer(rrs, dns.TypeA)
|
|
if len(ips) != 1 {
|
|
t.Fatalf("expected 1 A IP, got %d", len(ips))
|
|
}
|
|
})
|
|
}
|
|
|
|
func TestDiscoverRootsOverride(t *testing.T) {
|
|
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
|
defer cancel()
|
|
|
|
cfg := &RootDiscoveryConfig{
|
|
Server: "a.root-servers.net",
|
|
IncludeAAAA: false,
|
|
}
|
|
|
|
servers, err := DiscoverRoots(ctx, cfg)
|
|
if err != nil {
|
|
t.Logf("skipping (no resolver available): %v", err)
|
|
t.Skip()
|
|
}
|
|
if len(servers) == 0 {
|
|
t.Fatal("expected at least one root server")
|
|
}
|
|
if len(servers[0].IPv4) == 0 {
|
|
t.Error("expected IPv4 addresses for a.root-servers.net")
|
|
}
|
|
}
|
|
|
|
func TestDiscoverRootsSingle(t *testing.T) {
|
|
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
|
defer cancel()
|
|
|
|
servers, err := DiscoverRoots(ctx, nil)
|
|
if err != nil {
|
|
t.Logf("skipping (no resolver available): %v", err)
|
|
t.Skip()
|
|
}
|
|
if len(servers) == 0 {
|
|
t.Fatal("expected at least one root server")
|
|
}
|
|
if servers[0].Name == "" {
|
|
t.Error("root server name should not be empty")
|
|
}
|
|
}
|
|
|
|
func TestDiscoverRootsNilConfig(t *testing.T) {
|
|
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
|
defer cancel()
|
|
|
|
servers, err := DiscoverRoots(ctx, nil)
|
|
if err != nil {
|
|
t.Logf("skipping (no resolver available): %v", err)
|
|
t.Skip()
|
|
}
|
|
if len(servers) == 0 {
|
|
t.Fatal("expected at least one root server with nil config")
|
|
}
|
|
}
|
|
|
|
func TestBuildNSResponse(t *testing.T) {
|
|
msg := new(dns.Msg)
|
|
msg.SetReply(new(dns.Msg))
|
|
msg.Answer = append(msg.Answer,
|
|
&dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS, Class: dns.ClassINET}, Ns: "a.root-servers.net."},
|
|
&dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS, Class: dns.ClassINET}, Ns: "b.root-servers.net."},
|
|
&dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS, Class: dns.ClassINET}, Ns: "c.root-servers.net."},
|
|
)
|
|
|
|
names := extractNSRecords(msg.Answer)
|
|
if len(names) != 3 {
|
|
t.Fatalf("expected 3 NS records, got %d", len(names))
|
|
}
|
|
|
|
for _, name := range names {
|
|
if len(name) == 0 || name[len(name)-1] != '.' {
|
|
t.Errorf("expected FQDN, got %q", name)
|
|
}
|
|
}
|
|
}
|
|
|
|
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("extractNSNames: expected 2 names, got %d", len(names))
|
|
}
|
|
}
|
|
|
|
func TestExtractNSNamesEmpty(t *testing.T) {
|
|
names := extractNSNames(nil)
|
|
if len(names) != 0 {
|
|
t.Errorf("extractNSNames(nil): expected 0 names, got %d", len(names))
|
|
}
|
|
}
|
|
|
|
func TestExtractNSNamesNonNS(t *testing.T) {
|
|
rrs := []dns.RR{
|
|
&dns.A{Hdr: dns.RR_Header{Name: "a.root-servers.net.", Rrtype: dns.TypeA}, A: net.ParseIP("198.41.0.4")},
|
|
}
|
|
names := extractNSNames(rrs)
|
|
if len(names) != 0 {
|
|
t.Errorf("extractNSNames with A records: expected 0 names, got %d", len(names))
|
|
}
|
|
}
|
|
|
|
func TestDiscoverRootsAllRoots(t *testing.T) {
|
|
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
|
|
defer cancel()
|
|
|
|
cfg := &RootDiscoveryConfig{
|
|
AllRoots: true,
|
|
IncludeAAAA: false,
|
|
}
|
|
|
|
servers, err := DiscoverRoots(ctx, cfg)
|
|
if err != nil {
|
|
t.Logf("skipping (no resolver available): %v", err)
|
|
t.Skip()
|
|
}
|
|
if len(servers) == 0 {
|
|
t.Fatal("expected root servers with AllRoots=true")
|
|
}
|
|
t.Logf("discovered %d root servers", len(servers))
|
|
}
|
|
|
|
func TestDiscoverRootsAllRootsIncludeAAAA(t *testing.T) {
|
|
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
|
|
defer cancel()
|
|
|
|
cfg := &RootDiscoveryConfig{
|
|
AllRoots: true,
|
|
IncludeAAAA: true,
|
|
}
|
|
|
|
servers, err := DiscoverRoots(ctx, cfg)
|
|
if err != nil {
|
|
t.Logf("skipping (no resolver available): %v", err)
|
|
t.Skip()
|
|
}
|
|
if len(servers) == 0 {
|
|
t.Fatal("expected root servers")
|
|
}
|
|
}
|
|
|
|
// startMockDNSServer starts a UDP DNS server on a random port that serves
|
|
// pre-configured responses. It returns the server address and a stop function.
|
|
func startMockDNSServer(t *testing.T, handlerFn dns.HandlerFunc) string {
|
|
t.Helper()
|
|
|
|
mux := dns.NewServeMux()
|
|
mux.HandleFunc(".", handlerFn)
|
|
|
|
srv := &dns.Server{
|
|
Addr: "127.0.0.1:0",
|
|
Net: "udp",
|
|
Handler: mux,
|
|
}
|
|
|
|
started := make(chan struct{})
|
|
srv.NotifyStartedFunc = func() { close(started) }
|
|
|
|
go func() {
|
|
if err := srv.ListenAndServe(); err != nil && t.Failed() {
|
|
return
|
|
}
|
|
}()
|
|
|
|
select {
|
|
case <-started:
|
|
case <-time.After(2 * time.Second):
|
|
t.Fatal("mock DNS server did not start in time")
|
|
}
|
|
|
|
// Retrieve the actual bound address from the server's PacketConn.
|
|
addr := srv.PacketConn.LocalAddr().String()
|
|
t.Cleanup(func() { _ = srv.Shutdown() })
|
|
return addr
|
|
}
|
|
|
|
func TestQueryResolverSuccess(t *testing.T) {
|
|
addr := startMockDNSServer(t, func(w dns.ResponseWriter, r *dns.Msg) {
|
|
m := new(dns.Msg)
|
|
m.SetReply(r)
|
|
m.Answer = append(m.Answer,
|
|
&dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS, Class: dns.ClassINET, Ttl: 300}, Ns: "a.root-servers.net."},
|
|
)
|
|
_ = w.WriteMsg(m)
|
|
})
|
|
|
|
// queryResolver uses 5ns timeout without deadline; provide one
|
|
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
|
|
defer cancel()
|
|
msg, err := queryResolver(ctx, addr, ".", dns.TypeNS)
|
|
if err != nil {
|
|
t.Fatalf("queryResolver: %v", err)
|
|
}
|
|
names := extractNSRecords(msg.Answer)
|
|
if len(names) == 0 {
|
|
t.Fatal("expected NS records in answer")
|
|
}
|
|
}
|
|
|
|
func TestQueryResolverError(t *testing.T) {
|
|
ctx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond)
|
|
defer cancel()
|
|
// Use an address nothing is listening on
|
|
_, err := queryResolver(ctx, "127.0.0.1:19999", ".", dns.TypeNS)
|
|
if err == nil {
|
|
t.Fatal("expected error for unreachable resolver")
|
|
}
|
|
}
|
|
|
|
func TestDiscoverSingleRootWithMock(t *testing.T) {
|
|
addr := startMockDNSServer(t, func(w dns.ResponseWriter, r *dns.Msg) {
|
|
m := new(dns.Msg)
|
|
m.SetReply(r)
|
|
q := r.Question[0]
|
|
switch q.Qtype {
|
|
case dns.TypeNS:
|
|
m.Answer = append(m.Answer,
|
|
&dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS, Class: dns.ClassINET, Ttl: 300}, Ns: "mock.root-servers.test."},
|
|
)
|
|
case dns.TypeA:
|
|
m.Answer = append(m.Answer,
|
|
&dns.A{Hdr: dns.RR_Header{Name: q.Name, Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, A: net.ParseIP("127.0.0.1")},
|
|
)
|
|
}
|
|
_ = w.WriteMsg(m)
|
|
})
|
|
|
|
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
|
|
defer cancel()
|
|
msg, err := queryResolver(ctx, addr, ".", dns.TypeNS)
|
|
if err != nil {
|
|
t.Fatalf("queryResolver: %v", err)
|
|
}
|
|
names := extractNSRecords(msg.Answer)
|
|
if len(names) == 0 {
|
|
names = extractNSNames(msg.Ns)
|
|
}
|
|
if len(names) == 0 {
|
|
t.Skip("mock NS query returned no NS records")
|
|
}
|
|
t.Logf("found %d root NS names from mock: %v", len(names), names)
|
|
}
|
|
|
|
func TestDiscoverAllRootsWithMockServer(t *testing.T) {
|
|
addr := startMockDNSServer(t, func(w dns.ResponseWriter, r *dns.Msg) {
|
|
m := new(dns.Msg)
|
|
m.SetReply(r)
|
|
q := r.Question[0]
|
|
switch q.Qtype {
|
|
case dns.TypeNS:
|
|
m.Answer = append(m.Answer,
|
|
&dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS, Class: dns.ClassINET, Ttl: 300}, Ns: "a.mock-roots.test."},
|
|
&dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS, Class: dns.ClassINET, Ttl: 300}, Ns: "b.mock-roots.test."},
|
|
)
|
|
case dns.TypeA:
|
|
m.Answer = append(m.Answer,
|
|
&dns.A{Hdr: dns.RR_Header{Name: q.Name, Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, A: net.ParseIP("127.0.0.1")},
|
|
)
|
|
}
|
|
_ = w.WriteMsg(m)
|
|
})
|
|
|
|
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
|
defer cancel()
|
|
msg, err := queryResolver(ctx, addr, ".", dns.TypeNS)
|
|
if err != nil {
|
|
t.Fatalf("queryResolver: %v", err)
|
|
}
|
|
names := extractNSRecords(msg.Answer)
|
|
if len(names) < 2 {
|
|
t.Fatalf("expected 2 NS names, got %d", len(names))
|
|
}
|
|
|
|
// Also cover AAAA path
|
|
aaaaMsg, err := queryResolver(ctx, addr, "a.mock-roots.test.", dns.TypeAAAA)
|
|
if err != nil {
|
|
t.Logf("AAAA query error (acceptable): %v", err)
|
|
} else {
|
|
t.Logf("AAAA query returned %d answers", len(aaaaMsg.Answer))
|
|
}
|
|
}
|
|
|
|
func TestMinTTLFromMsgWithExtraRecords(t *testing.T) {
|
|
msg := new(dns.Msg)
|
|
msg.Answer = append(msg.Answer, &dns.A{
|
|
Hdr: dns.RR_Header{Ttl: 300},
|
|
A: net.ParseIP("1.2.3.4"),
|
|
})
|
|
// Extra record (non-OPT) with smaller TTL
|
|
msg.Extra = append(msg.Extra, &dns.NS{
|
|
Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS, Ttl: 60},
|
|
Ns: "a.root-servers.net.",
|
|
})
|
|
|
|
ttl := minTTLFromMsg(msg)
|
|
if ttl != 60*time.Second {
|
|
t.Errorf("minTTLFromMsg = %v, want 60s", ttl)
|
|
}
|
|
}
|
|
|
|
func TestMinTTLFromMsgOPTIgnored(t *testing.T) {
|
|
msg := new(dns.Msg)
|
|
msg.Answer = append(msg.Answer, &dns.A{
|
|
Hdr: dns.RR_Header{Ttl: 300},
|
|
A: net.ParseIP("1.2.3.4"),
|
|
})
|
|
// OPT record should be ignored
|
|
msg.Extra = append(msg.Extra, &dns.OPT{
|
|
Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeOPT},
|
|
})
|
|
|
|
ttl := minTTLFromMsg(msg)
|
|
if ttl != 300*time.Second {
|
|
t.Errorf("minTTLFromMsg with OPT = %v, want 300s", ttl)
|
|
}
|
|
}
|