265 lines
6.6 KiB
Go
265 lines
6.6 KiB
Go
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), ".")
|
|
}
|