Add RootServer struct and DiscoverRoots function that queries the local resolver for NS records of the root zone, resolves each root name to A/AAAA addresses, and returns a list of RootServer structs. Supports: - Default: discovers a single root server from local resolver - --root-server override to target a specific root server - --all-root-servers to discover all 13 root servers - --root-aaaa flag to include IPv6 addresses Includes unit tests for all helper functions and integration tests that skip gracefully when no resolver is available. Co-authored-by: multica-agent <github@multica.ai>
This commit is contained in:
@@ -0,0 +1,227 @@
|
||||
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)
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user