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