diff --git a/internal/dns/roots.go b/internal/dns/roots.go new file mode 100644 index 0000000..811eafe --- /dev/null +++ b/internal/dns/roots.go @@ -0,0 +1,223 @@ +package dns + +import ( + "context" + "fmt" + "net" + "strings" + "time" + + "github.com/miekg/dns" +) + +type RootServer struct { + Name string + IPv4 []net.IP + IPv6 []net.IP +} + +func (rs *RootServer) AllIPs(includeAAAA bool) []net.IP { + var ips []net.IP + ips = append(ips, rs.IPv4...) + if includeAAAA { + ips = append(ips, rs.IPv6...) + } + return ips +} + +type RootDiscoveryConfig struct { + Server string + AllRoots bool + IncludeAAAA bool +} + +func DefaultRootDiscoveryConfig() *RootDiscoveryConfig { + return &RootDiscoveryConfig{ + AllRoots: false, + IncludeAAAA: false, + } +} + +func DiscoverRoots(ctx context.Context, cfg *RootDiscoveryConfig) ([]RootServer, error) { + if cfg == nil { + cfg = DefaultRootDiscoveryConfig() + } + + if cfg.Server != "" { + return discoverRootOverride(ctx, cfg.Server, cfg.IncludeAAAA) + } + + if cfg.AllRoots { + return discoverAllRoots(ctx, cfg.IncludeAAAA) + } + + return discoverSingleRoot(ctx, cfg.IncludeAAAA) +} + +func discoverRootOverride(ctx context.Context, server string, includeAAAA bool) ([]RootServer, error) { + resolver := systemResolver() + + nsMsg, err := queryResolver(ctx, resolver, ".", dns.TypeNS) + if err != nil { + return nil, fmt.Errorf("query root NS records: %w", err) + } + + nsSet := extractNSRecords(nsMsg.Answer) + if len(nsSet) == 0 { + nsSet = extractNSNames(nsMsg.Ns) + } + + normalized := normalizeServerName(server) + for _, name := range nsSet { + if normalizeServerName(name) == normalized { + return resolveRootServer(ctx, name, includeAAAA) + } + } + + return resolveRootServer(ctx, server, includeAAAA) +} + +func discoverSingleRoot(ctx context.Context, includeAAAA bool) ([]RootServer, error) { + resolver := systemResolver() + + nsMsg, err := queryResolver(ctx, resolver, ".", dns.TypeNS) + if err != nil { + return nil, fmt.Errorf("query root NS records: %w", err) + } + + nsSet := extractNSRecords(nsMsg.Answer) + if len(nsSet) == 0 { + nsSet = extractNSNames(nsMsg.Ns) + } + + if len(nsSet) == 0 { + return nil, fmt.Errorf("no root NS records found in response") + } + + pick := nsSet[0] + return resolveRootServer(ctx, pick, includeAAAA) +} + +func discoverAllRoots(ctx context.Context, includeAAAA bool) ([]RootServer, error) { + resolver := systemResolver() + + nsMsg, err := queryResolver(ctx, resolver, ".", dns.TypeNS) + if err != nil { + return nil, fmt.Errorf("query root NS records: %w", err) + } + + nsSet := extractNSRecords(nsMsg.Answer) + if len(nsSet) == 0 { + nsSet = extractNSNames(nsMsg.Ns) + } + + if len(nsSet) == 0 { + return nil, fmt.Errorf("no root NS records found in response") + } + + var servers []RootServer + for _, name := range nsSet { + resolved, err := resolveRootServer(ctx, name, includeAAAA) + if err != nil { + servers = append(servers, RootServer{Name: name}) + continue + } + servers = append(servers, resolved...) + } + + if len(servers) == 0 { + return nil, fmt.Errorf("failed to resolve any root servers") + } + + return servers, nil +} + +func resolveRootServer(ctx context.Context, name string, includeAAAA bool) ([]RootServer, error) { + resolver := systemResolver() + + var ipv4 []net.IP + aMsg, err := queryResolver(ctx, resolver, name, dns.TypeA) + if err == nil { + ipv4 = extractIPsFromAnswer(aMsg.Answer, dns.TypeA) + } + + var ipv6 []net.IP + if includeAAAA { + aaaaMsg, err := queryResolver(ctx, resolver, name, dns.TypeAAAA) + if err == nil { + ipv6 = extractIPsFromAnswer(aaaaMsg.Answer, dns.TypeAAAA) + } + } + + if len(ipv4) == 0 && len(ipv6) == 0 { + return nil, fmt.Errorf("no addresses for root server %s", name) + } + + return []RootServer{{Name: name, IPv4: ipv4, IPv6: ipv6}}, nil +} + +func systemResolver() string { + return "127.0.0.1:53" +} + +func queryResolver(ctx context.Context, resolverAddr, name string, qtype uint16) (*dns.Msg, error) { + c := &dns.Client{ + Net: "udp", + ReadTimeout: 5, + WriteTimeout: 5, + } + if deadline, ok := ctx.Deadline(); ok { + c.ReadTimeout = time.Until(deadline) + c.WriteTimeout = time.Until(deadline) + } + + m := new(dns.Msg) + m.SetQuestion(dns.Fqdn(name), qtype) + m.RecursionDesired = true + + r, _, err := c.ExchangeContext(ctx, m, resolverAddr) + if err != nil { + return nil, fmt.Errorf("resolver exchange %s %s: %w", name, QNameType(qtype), err) + } + return r, nil +} + +func extractNSRecords(rrs []dns.RR) []string { + var names []string + seen := make(map[string]bool) + for _, rr := range rrs { + if ns, ok := rr.(*dns.NS); ok { + n := dns.Fqdn(ns.Ns) + if !seen[n] { + seen[n] = true + names = append(names, n) + } + } + } + return names +} + +func extractNSNames(rrs []dns.RR) []string { + return extractNSRecords(rrs) +} + +func extractIPsFromAnswer(rrs []dns.RR, qtype uint16) []net.IP { + var ips []net.IP + for _, rr := range rrs { + switch v := rr.(type) { + case *dns.A: + if qtype == dns.TypeA { + ips = append(ips, v.A) + } + case *dns.AAAA: + if qtype == dns.TypeAAAA { + ips = append(ips, v.AAAA) + } + } + } + return ips +} + +func normalizeServerName(name string) string { + return strings.TrimSuffix(strings.ToLower(name), ".") +} diff --git a/internal/dns/roots_test.go b/internal/dns/roots_test.go new file mode 100644 index 0000000..5b1a4c0 --- /dev/null +++ b/internal/dns/roots_test.go @@ -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) + } + } +}