Add RootServer struct and DiscoverRoots function that queries the local resolver for NS records of the root zone, resolves each root name to A/AAAA addresses, and returns a list of RootServer structs. Supports: - Default: discovers a single root server from local resolver - --root-server override to target a specific root server - --all-root-servers to discover all 13 root servers - --root-aaaa flag to include IPv6 addresses Includes unit tests for all helper functions and integration tests that skip gracefully when no resolver is available. Co-authored-by: multica-agent <github@multica.ai>
This commit is contained in:
@@ -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), ".")
|
||||
}
|
||||
Reference in New Issue
Block a user