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 overrides which root server to use as the traversal starting point. // When empty, a root server is discovered via the upstream resolver. Server string // Resolver is the upstream DNS resolver used to resolve root server names. // When empty, the system resolver from /etc/resolv.conf is used. // Format: "host:port" (e.g. "8.8.8.8:53" or "1.1.1.1:53"). Resolver 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() } resolver := resolverFromConfig(cfg) if cfg.Server != "" { return discoverRootOverride(ctx, resolver, cfg.Server, cfg.IncludeAAAA) } if cfg.AllRoots { servers, err := discoverAllRoots(ctx, resolver, cfg.IncludeAAAA) if err != nil { return filterHints(RootHints, cfg.IncludeAAAA), nil } return servers, nil } servers, err := discoverSingleRoot(ctx, resolver, cfg.IncludeAAAA) if err != nil { hints := filterHints(RootHints, cfg.IncludeAAAA) if len(hints) > 0 { return hints[:1], nil } return nil, err } return servers, nil } // filterHints returns a copy of hints with IPv6 addresses stripped when includeAAAA is false. func filterHints(hints []RootServer, includeAAAA bool) []RootServer { out := make([]RootServer, len(hints)) for i, h := range hints { out[i] = RootServer{Name: h.Name, IPv4: h.IPv4} if includeAAAA { out[i].IPv6 = h.IPv6 } } return out } func discoverRootOverride(ctx context.Context, resolver, server string, includeAAAA bool) ([]RootServer, error) { 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, resolver, name, includeAAAA) } } return resolveRootServer(ctx, resolver, server, includeAAAA) } func discoverSingleRoot(ctx context.Context, resolver string, includeAAAA bool) ([]RootServer, error) { 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, resolver, pick, includeAAAA) } func discoverAllRoots(ctx context.Context, resolver string, includeAAAA bool) ([]RootServer, error) { 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, resolver, 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, resolver, name string, includeAAAA bool) ([]RootServer, error) { 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 } // resolverFromConfig returns the upstream DNS resolver address to use. // If cfg.Resolver is set, it is used directly. Otherwise the system resolver // is read from /etc/resolv.conf. Falls back to 127.0.0.1:53 if neither is available. func resolverFromConfig(cfg *RootDiscoveryConfig) string { if cfg != nil && cfg.Resolver != "" { return cfg.Resolver } return systemResolver() } // systemResolver returns the first nameserver from the system DNS configuration. // This is Unix-only: it reads /etc/resolv.conf, which does not exist on Windows. // On Windows (or any system without /etc/resolv.conf) the fallback 127.0.0.1:53 applies. func systemResolver() string { cc, err := dns.ClientConfigFromFile("/etc/resolv.conf") if err != nil || len(cc.Servers) == 0 { return "127.0.0.1:53" } return net.JoinHostPort(cc.Servers[0], cc.Port) } func queryResolver(ctx context.Context, resolverAddr, name string, qtype uint16) (*dns.Msg, error) { c := &dns.Client{ Net: "udp", ReadTimeout: 5 * time.Second, WriteTimeout: 5 * time.Second, } 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), ".") }