package dns import ( "context" "errors" "fmt" "net" "strings" "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. It accepts a hostname (resolved to A via the upstream resolver) // or an IP literal (used directly, no lookup). Server string // Resolver is the upstream DNS resolver used to discover and resolve // root server names. When empty, the system resolver configuration is // used. Format: "host:port" (e.g. "8.8.8.8:53" or "1.1.1.1:53"). Resolver string AllRoots bool IncludeAAAA bool // Query controls transport parameters (retries, timeout, TCP fallback) // for discovery queries. nil means DefaultQueryConfig. Query *QueryConfig // Exchange overrides the wire exchange; nil means the real network. // Tests inject a mock here so discovery is network-free. Exchange ExchangeFunc } func DefaultRootDiscoveryConfig() *RootDiscoveryConfig { return &RootDiscoveryConfig{ AllRoots: false, IncludeAAAA: false, } } func DiscoverRoots(ctx context.Context, cfg *RootDiscoveryConfig) ([]RootServer, error) { if cfg == nil { cfg = DefaultRootDiscoveryConfig() } // --root-server with an IP literal: use it directly, never look it up. if cfg.Server != "" { if ip := net.ParseIP(cfg.Server); ip != nil { rs := RootServer{Name: cfg.Server} if ip.To4() != nil { rs.IPv4 = []net.IP{ip} } else { rs.IPv6 = []net.IP{ip} } return []RootServer{rs}, nil } } resolver, err := resolverFromConfig(cfg) if err != nil { if cfg.Server != "" { return nil, err } return rootHintsFallback(cfg, err) } if cfg.Server != "" { return resolveRootServer(ctx, cfg, resolver, cfg.Server) } if cfg.AllRoots { servers, err := discoverAllRoots(ctx, cfg, resolver) if err != nil { return rootHintsFallback(cfg, err) } return servers, nil } servers, err := discoverSingleRoot(ctx, cfg, resolver) if err != nil { return rootHintsFallback(cfg, err) } return servers, nil } // rootHintsFallback returns the builtin IANA hints when the upstream resolver // cannot provide roots (one entry unless AllRoots). func rootHintsFallback(cfg *RootDiscoveryConfig, err error) ([]RootServer, error) { hints := filterHints(RootHints, cfg.IncludeAAAA) if len(hints) == 0 { return nil, err } if cfg.AllRoots { return hints, nil } return hints[:1], 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 } // discoverSingleRoot mirrors get_a_root in traverser.rb: ask the upstream // resolver for the root NS set, prefer glue from the additional section, and // only fall back to explicit A/AAAA lookups when no glue was supplied. func discoverSingleRoot(ctx context.Context, cfg *RootDiscoveryConfig, resolver string) ([]RootServer, error) { names, nsMsg, err := rootNSNames(ctx, cfg, resolver) if err != nil { return nil, err } for _, name := range names { rs := rootFromAdditional(nsMsg, name, cfg.IncludeAAAA) if len(rs.AllIPs(cfg.IncludeAAAA)) > 0 { return []RootServer{singleAddress(rs)}, nil } } var lastErr error for _, name := range names { servers, err := resolveRootServer(ctx, cfg, resolver, name) if err != nil { lastErr = err continue } return []RootServer{singleAddress(servers[0])}, nil } if lastErr == nil { lastErr = errors.New("no address could be found for any root server") } return nil, lastErr } // singleAddress narrows rs to its first address, mirroring get_a_root in // traverser.rb (add[0]/ans2[0]): the single-root start point is exactly one // (name, IP) pair even when the upstream supplies more (or duplicate) glue. func singleAddress(rs RootServer) RootServer { if len(rs.IPv4) > 0 { return RootServer{Name: rs.Name, IPv4: rs.IPv4[:1]} } if len(rs.IPv6) > 0 { return RootServer{Name: rs.Name, IPv6: rs.IPv6[:1]} } return rs } // discoverAllRoots mirrors find_all_roots in traverser.rb: it returns one // RootServer per root NS name with its full address set, so the traversal can // branch per root. Roots with no resolvable address are skipped. func discoverAllRoots(ctx context.Context, cfg *RootDiscoveryConfig, resolver string) ([]RootServer, error) { names, nsMsg, err := rootNSNames(ctx, cfg, resolver) if err != nil { return nil, err } var servers []RootServer for _, name := range names { rs := rootFromAdditional(nsMsg, name, cfg.IncludeAAAA) if len(rs.AllIPs(cfg.IncludeAAAA)) == 0 { resolved, err := resolveRootServer(ctx, cfg, resolver, name) if err != nil { continue } rs = resolved[0] } servers = append(servers, rs) } if len(servers) == 0 { return nil, errors.New("failed to resolve any root servers") } return servers, nil } func rootNSNames(ctx context.Context, cfg *RootDiscoveryConfig, resolver string) ([]string, *dns.Msg, error) { nsMsg, err := queryUpstream(ctx, cfg, resolver, ".", dns.TypeNS) if err != nil { return nil, nil, fmt.Errorf("query root NS records: %w", err) } names := extractNSRecords(nsMsg.Answer) if len(names) == 0 { names = extractNSNames(nsMsg.Ns) } if len(names) == 0 { return nil, nil, errors.New("no root NS records found in response") } return names, nsMsg, nil } func rootFromAdditional(msg *dns.Msg, name string, includeAAAA bool) RootServer { rs := RootServer{Name: name, IPv4: additionalIPs(msg, name, dns.TypeA)} if includeAAAA { rs.IPv6 = additionalIPs(msg, name, dns.TypeAAAA) } return rs } func resolveRootServer(ctx context.Context, cfg *RootDiscoveryConfig, resolver, name string) ([]RootServer, error) { var ipv4 []net.IP aMsg, err := queryUpstream(ctx, cfg, resolver, name, dns.TypeA) if err == nil { ipv4 = extractIPsFromAnswer(aMsg.Answer, dns.TypeA) } var ipv6 []net.IP if cfg.IncludeAAAA { aaaaMsg, err := queryUpstream(ctx, cfg, 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: // cfg.Resolver when set, otherwise the system resolver configuration. There // is deliberately no hardcoded address fallback. func resolverFromConfig(cfg *RootDiscoveryConfig) (string, error) { if cfg != nil && cfg.Resolver != "" { return cfg.Resolver, nil } 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; there the caller falls back to the builtin root hints. func systemResolver() (string, error) { cc, err := dns.ClientConfigFromFile("/etc/resolv.conf") if err != nil { return "", fmt.Errorf("read system resolver config: %w", err) } if len(cc.Servers) == 0 { return "", errors.New("no nameservers found in system resolver config") } return net.JoinHostPort(cc.Servers[0], cc.Port), nil } // queryUpstream asks the upstream resolver with recursion desired — the only // RD=1 path in the program — using the same retry/TCP-fallback machinery as // traversal queries. func queryUpstream(ctx context.Context, cfg *RootDiscoveryConfig, resolverAddr, name string, qtype uint16) (*dns.Msg, error) { var qcfg *QueryConfig var exchange ExchangeFunc if cfg != nil { qcfg = cfg.Query exchange = cfg.Exchange } qcfg = qcfg.withDefaults() if exchange == nil { exchange = realExchange } m := new(dns.Msg) m.SetQuestion(dns.Fqdn(name), qtype) m.RecursionDesired = true if qcfg.UDPSize > MinEDNS0UDPSize() { m.SetEdns0(uint16(qcfg.UDPSize), false) } return exchangeWithRetry(ctx, exchange, resolverAddr, m, qcfg) } 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 } // additionalIPs returns glue addresses for name from the additional section. func additionalIPs(msg *dns.Msg, name string, qtype uint16) []net.IP { var ips []net.IP for _, rr := range msg.Extra { if !strings.EqualFold(rr.Header().Name, name) { continue } 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 }