package dns import ( "context" "fmt" "net" "strings" "time" "github.com/miekg/dns" ) // RootServer holds the name and IP addresses of a DNS root nameserver. type RootServer struct { // Name is the FQDN of the root nameserver (e.g. "a.root-servers.net."). Name string // IPv4 holds the IPv4 addresses for the server. IPv4 []net.IP // IPv6 holds the IPv6 addresses for the server. IPv6 []net.IP } // AllIPs returns all IP addresses for the server. // When includeAAAA is false, only IPv4 addresses are returned. 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 } // RootDiscoveryConfig controls how DiscoverRoots selects root servers. type RootDiscoveryConfig struct { // Server overrides which root server is used. An empty string means // auto-select the first root server returned by the system resolver. Server string // AllRoots queries all 13 root servers instead of just one. AllRoots bool // IncludeAAAA includes IPv6 addresses of root servers when true. IncludeAAAA bool } // DefaultRootDiscoveryConfig returns a RootDiscoveryConfig that auto-selects a // single IPv4-only root server. func DefaultRootDiscoveryConfig() *RootDiscoveryConfig { return &RootDiscoveryConfig{ AllRoots: false, IncludeAAAA: false, } } // DiscoverRoots discovers DNS root servers to use as traversal starting points. // When cfg.Server is set, that specific root server is used. // When cfg.AllRoots is true, all 13 root servers are returned. // Otherwise, a single root server is selected from the system resolver's NS response. 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), ".") }