package dns import ( "context" "errors" "net" "strings" "testing" "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."} if len(empty.AllIPs(false)) != 0 { t.Error("expected 0 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") } } // mockUpstream builds an ExchangeFunc that answers root discovery queries. // glue controls whether A records are put in the additional section of the // NS response. func mockUpstream(t *testing.T, roots map[string]string, glue bool) ExchangeFunc { t.Helper() return func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { if !msg.RecursionDesired { t.Error("root discovery must query the upstream resolver with RD=1") } q := msg.Question[0] m := new(dns.Msg) m.SetReply(msg) switch { case q.Name == "." && q.Qtype == dns.TypeNS: for name, ip := range roots { m.Answer = append(m.Answer, &dns.NS{ Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS, Class: dns.ClassINET, Ttl: 300}, Ns: name, }) if glue { m.Extra = append(m.Extra, &dns.A{ Hdr: dns.RR_Header{Name: name, Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, A: net.ParseIP(ip), }) } } case q.Qtype == dns.TypeA: if ip, ok := roots[q.Name]; ok { 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(ip), }) } } return m, nil } } func TestDiscoverRootsIPLiteral(t *testing.T) { cfg := &RootDiscoveryConfig{ Server: "192.203.230.10", Exchange: func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { t.Error("IP literal root must not trigger any lookup") return nil, errors.New("no network") }, } servers, err := DiscoverRoots(context.Background(), cfg) if err != nil { t.Fatalf("DiscoverRoots: %v", err) } if len(servers) != 1 { t.Fatalf("expected 1 root, got %d", len(servers)) } if servers[0].Name != "192.203.230.10" { t.Errorf("Name = %q, want the IP literal", servers[0].Name) } if len(servers[0].IPv4) != 1 || !servers[0].IPv4[0].Equal(net.ParseIP("192.203.230.10")) { t.Errorf("IPv4 = %v, want [192.203.230.10]", servers[0].IPv4) } } func TestDiscoverRootsIPv6Literal(t *testing.T) { cfg := &RootDiscoveryConfig{Server: "2001:503:ba3e::2:30"} servers, err := DiscoverRoots(context.Background(), cfg) if err != nil { t.Fatalf("DiscoverRoots: %v", err) } if len(servers) != 1 || len(servers[0].IPv6) != 1 { t.Fatalf("expected 1 root with 1 IPv6 address, got %+v", servers) } } func TestDiscoverRootsHostnameOverride(t *testing.T) { roots := map[string]string{"e.root-servers.net.": "192.203.230.10"} cfg := &RootDiscoveryConfig{ Server: "e.root-servers.net", Resolver: "192.0.2.53:53", Exchange: mockUpstream(t, roots, false), } servers, err := DiscoverRoots(context.Background(), cfg) if err != nil { t.Fatalf("DiscoverRoots: %v", err) } if len(servers) != 1 { t.Fatalf("expected 1 root, got %d", len(servers)) } if servers[0].Name != "e.root-servers.net" { t.Errorf("Name = %q", servers[0].Name) } if len(servers[0].IPv4) != 1 || !servers[0].IPv4[0].Equal(net.ParseIP("192.203.230.10")) { t.Errorf("IPv4 = %v", servers[0].IPv4) } } func TestDiscoverRootsHostnameOverrideUnresolvable(t *testing.T) { cfg := &RootDiscoveryConfig{ Server: "nonexistent.root-servers.net", Resolver: "192.0.2.53:53", Exchange: mockUpstream(t, map[string]string{}, false), } if _, err := DiscoverRoots(context.Background(), cfg); err == nil { t.Fatal("expected error for unresolvable root override") } } func TestDiscoverRootsSingleFromGlue(t *testing.T) { roots := map[string]string{"a.root-servers.net.": "198.41.0.4"} var wireQueries int inner := mockUpstream(t, roots, true) cfg := &RootDiscoveryConfig{ Resolver: "192.0.2.53:53", Exchange: func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { wireQueries++ return inner(ctx, server, msg, useTCP) }, } servers, err := DiscoverRoots(context.Background(), cfg) if err != nil { t.Fatalf("DiscoverRoots: %v", err) } if len(servers) != 1 { t.Fatalf("expected exactly one root, got %d", len(servers)) } if len(servers[0].IPv4) == 0 { t.Error("expected glue A address") } if wireQueries != 1 { t.Errorf("glue should satisfy discovery in one query, got %d", wireQueries) } } func TestDiscoverRootsSingleWithoutGlue(t *testing.T) { roots := map[string]string{"a.root-servers.net.": "198.41.0.4"} cfg := &RootDiscoveryConfig{ Resolver: "192.0.2.53:53", Exchange: mockUpstream(t, roots, false), } servers, err := DiscoverRoots(context.Background(), cfg) if err != nil { t.Fatalf("DiscoverRoots: %v", err) } if len(servers) != 1 || len(servers[0].IPv4) != 1 { t.Fatalf("expected one root resolved via A lookup, got %+v", servers) } } func TestDiscoverAllRootsPerRootSet(t *testing.T) { roots := map[string]string{ "a.mock-roots.test.": "192.0.2.1", "b.mock-roots.test.": "192.0.2.2", "c.mock-roots.test.": "192.0.2.3", } cfg := &RootDiscoveryConfig{ AllRoots: true, Resolver: "192.0.2.53:53", Exchange: mockUpstream(t, roots, true), } servers, err := DiscoverRoots(context.Background(), cfg) if err != nil { t.Fatalf("DiscoverRoots: %v", err) } if len(servers) != 3 { t.Fatalf("expected 3 roots, got %d", len(servers)) } seen := map[string]bool{} for _, rs := range servers { seen[rs.Name] = true want, ok := roots[rs.Name] if !ok { t.Errorf("unexpected root %q", rs.Name) continue } if len(rs.IPv4) != 1 || !rs.IPv4[0].Equal(net.ParseIP(want)) { t.Errorf("root %s IPv4 = %v, want [%s]", rs.Name, rs.IPv4, want) } } if len(seen) != 3 { t.Errorf("roots not distinct: %v", seen) } } func TestDiscoverAllRootsSkipsUnresolvable(t *testing.T) { // b has neither glue nor an A record: it must be skipped, like // find_all_roots in traverser.rb. roots := map[string]string{"a.mock-roots.test.": "192.0.2.1"} exchange := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { q := msg.Question[0] m := new(dns.Msg) m.SetReply(msg) switch { case q.Name == "." && q.Qtype == dns.TypeNS: for _, name := range []string{"a.mock-roots.test.", "b.mock-roots.test."} { m.Answer = append(m.Answer, &dns.NS{ Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS, Class: dns.ClassINET, Ttl: 300}, Ns: name, }) } case q.Qtype == dns.TypeA: if ip, ok := roots[q.Name]; ok { 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(ip), }) } } return m, nil } cfg := &RootDiscoveryConfig{ AllRoots: true, Resolver: "192.0.2.53:53", Query: &QueryConfig{Retries: 1}, Exchange: exchange, } servers, err := DiscoverRoots(context.Background(), cfg) if err != nil { t.Fatalf("DiscoverRoots: %v", err) } if len(servers) != 1 || servers[0].Name != "a.mock-roots.test." { t.Fatalf("expected only the resolvable root, got %+v", servers) } } func TestDiscoverRootsFallsBackToHints(t *testing.T) { cfg := &RootDiscoveryConfig{ Resolver: "192.0.2.53:53", Query: &QueryConfig{Retries: 1}, Exchange: func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { return nil, errors.New("upstream unreachable") }, } servers, err := DiscoverRoots(context.Background(), cfg) if err != nil { t.Fatalf("expected hints fallback, got error: %v", err) } if len(servers) != 1 { t.Fatalf("single-root mode must fall back to one hint, got %d", len(servers)) } if !strings.HasSuffix(servers[0].Name, ".root-servers.net.") { t.Errorf("expected an IANA hint, got %q", servers[0].Name) } cfg.AllRoots = true servers, err = DiscoverRoots(context.Background(), cfg) if err != nil { t.Fatalf("expected hints fallback, got error: %v", err) } if len(servers) != len(RootHints) { t.Errorf("all-roots fallback should return all %d hints, got %d", len(RootHints), len(servers)) } } func TestDiscoverRootsRetriesUpstream(t *testing.T) { roots := map[string]string{"a.root-servers.net.": "198.41.0.4"} inner := mockUpstream(t, roots, true) var calls int cfg := &RootDiscoveryConfig{ Resolver: "192.0.2.53:53", Query: &QueryConfig{Retries: 2, RetryDelay: 1}, Exchange: func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { calls++ if calls == 1 { return nil, errors.New("transient failure") } return inner(ctx, server, msg, useTCP) }, } servers, err := DiscoverRoots(context.Background(), cfg) if err != nil { t.Fatalf("DiscoverRoots should retry: %v", err) } if len(servers) != 1 { t.Fatalf("expected 1 root after retry, got %d", len(servers)) } if calls != 2 { t.Errorf("expected 2 attempts, got %d", calls) } } func TestDiscoverRootsTruncationTCPFallback(t *testing.T) { roots := map[string]string{"a.root-servers.net.": "198.41.0.4"} inner := mockUpstream(t, roots, true) var sawTCP bool cfg := &RootDiscoveryConfig{ Resolver: "192.0.2.53:53", Query: &QueryConfig{Retries: 1, AllowTCP: true}, Exchange: func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { if !useTCP { m := new(dns.Msg) m.SetReply(msg) m.Truncated = true return m, nil } sawTCP = true return inner(ctx, server, msg, useTCP) }, } servers, err := DiscoverRoots(context.Background(), cfg) if err != nil { t.Fatalf("DiscoverRoots: %v", err) } if !sawTCP { t.Error("expected TCP fallback on truncated upstream response") } if len(servers) != 1 || len(servers[0].IPv4) == 0 { t.Fatalf("expected root from TCP response, got %+v", servers) } } func TestExtractNSRecords(t *testing.T) { t.Run("empty", func(t *testing.T) { if names := extractNSRecords(nil); 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."}, } if names := extractNSRecords(rrs); len(names) != 1 { t.Errorf("expected 1 deduped name, got %d", len(names)) } }) t.Run("non-NS ignored", func(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")}, } if names := extractNSNames(rrs); len(names) != 0 { t.Errorf("expected 0 names, got %d", len(names)) } }) } func TestExtractIPsFromAnswer(t *testing.T) { t.Run("empty", func(t *testing.T) { if ips := extractIPsFromAnswer(nil, dns.TypeA); 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")}, } if ips := extractIPsFromAnswer(rrs, dns.TypeA); len(ips) != 2 { t.Fatalf("expected 2 IPs, 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")}, } if ips := extractIPsFromAnswer(rrs, dns.TypeA); len(ips) != 1 { t.Fatalf("expected 1 A IP, got %d", len(ips)) } if ips := extractIPsFromAnswer(rrs, dns.TypeAAAA); len(ips) != 1 { t.Fatalf("expected 1 AAAA IP, got %d", len(ips)) } }) } func TestAdditionalIPs(t *testing.T) { msg := new(dns.Msg) msg.Extra = append(msg.Extra, &dns.A{Hdr: dns.RR_Header{Name: "A.ROOT-SERVERS.NET.", Rrtype: dns.TypeA, Class: dns.ClassINET}, A: net.ParseIP("198.41.0.4")}, &dns.A{Hdr: dns.RR_Header{Name: "b.root-servers.net.", Rrtype: dns.TypeA, Class: dns.ClassINET}, A: net.ParseIP("170.247.170.2")}, &dns.AAAA{Hdr: dns.RR_Header{Name: "a.root-servers.net.", Rrtype: dns.TypeAAAA, Class: dns.ClassINET}, AAAA: net.ParseIP("2001:503:ba3e::2:30")}, ) ips := additionalIPs(msg, "a.root-servers.net.", dns.TypeA) if len(ips) != 1 || !ips[0].Equal(net.ParseIP("198.41.0.4")) { t.Errorf("case-insensitive glue match failed: %v", ips) } if ips := additionalIPs(msg, "a.root-servers.net.", dns.TypeAAAA); len(ips) != 1 { t.Errorf("expected 1 AAAA glue, got %v", ips) } if ips := additionalIPs(msg, "c.root-servers.net.", dns.TypeA); len(ips) != 0 { t.Errorf("expected no glue for c, got %v", ips) } }