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