feat: implement root server discovery (HAN-378) #4
@@ -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), ".")
|
||||||
|
}
|
||||||
@@ -0,0 +1,227 @@
|
|||||||
|
package dns
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"net"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/miekg/dns"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestRootServerAllIPs(t *testing.T) {
|
||||||
|
rs := RootServer{
|
||||||
|
Name: "a.root-servers.net.",
|
||||||
|
IPv4: []net.IP{net.ParseIP("198.41.0.4")},
|
||||||
|
IPv6: []net.IP{net.ParseIP("2001:503:ba3e::2:30")},
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Run("IPv4 only", func(t *testing.T) {
|
||||||
|
ips := rs.AllIPs(false)
|
||||||
|
if len(ips) != 1 {
|
||||||
|
t.Fatalf("expected 1 IP, got %d", len(ips))
|
||||||
|
}
|
||||||
|
if !ips[0].Equal(net.ParseIP("198.41.0.4")) {
|
||||||
|
t.Errorf("got %v, want 198.41.0.4", ips[0])
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("IPv4 and IPv6", func(t *testing.T) {
|
||||||
|
ips := rs.AllIPs(true)
|
||||||
|
if len(ips) != 2 {
|
||||||
|
t.Fatalf("expected 2 IPs, got %d", len(ips))
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("no addresses", func(t *testing.T) {
|
||||||
|
empty := RootServer{Name: "empty.root-servers.net."}
|
||||||
|
ips := empty.AllIPs(false)
|
||||||
|
if len(ips) != 0 {
|
||||||
|
t.Errorf("expected 0 IPs, got %d", len(ips))
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDefaultRootDiscoveryConfig(t *testing.T) {
|
||||||
|
cfg := DefaultRootDiscoveryConfig()
|
||||||
|
if cfg.AllRoots {
|
||||||
|
t.Error("AllRoots should be false by default")
|
||||||
|
}
|
||||||
|
if cfg.IncludeAAAA {
|
||||||
|
t.Error("IncludeAAAA should be false by default")
|
||||||
|
}
|
||||||
|
if cfg.Server != "" {
|
||||||
|
t.Error("Server should be empty by default")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNormalizeServerName(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
input string
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{"a.root-servers.net.", "a.root-servers.net"},
|
||||||
|
{"A.ROOT-SERVERS.NET.", "a.root-servers.net"},
|
||||||
|
{"b.root-servers.net", "b.root-servers.net"},
|
||||||
|
{"root-servers.net.", "root-servers.net"},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.input, func(t *testing.T) {
|
||||||
|
got := normalizeServerName(tt.input)
|
||||||
|
if got != tt.want {
|
||||||
|
t.Errorf("normalizeServerName(%q) = %q, want %q", tt.input, got, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExtractNSRecords(t *testing.T) {
|
||||||
|
t.Run("empty", func(t *testing.T) {
|
||||||
|
names := extractNSRecords(nil)
|
||||||
|
if len(names) != 0 {
|
||||||
|
t.Errorf("expected 0 names, got %d", len(names))
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("NS records", func(t *testing.T) {
|
||||||
|
rrs := []dns.RR{
|
||||||
|
&dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS}, Ns: "a.root-servers.net."},
|
||||||
|
&dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS}, Ns: "b.root-servers.net."},
|
||||||
|
}
|
||||||
|
names := extractNSRecords(rrs)
|
||||||
|
if len(names) != 2 {
|
||||||
|
t.Fatalf("expected 2 names, got %d", len(names))
|
||||||
|
}
|
||||||
|
if names[0] != "a.root-servers.net." {
|
||||||
|
t.Errorf("got %q, want a.root-servers.net.", names[0])
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("dedup", func(t *testing.T) {
|
||||||
|
rrs := []dns.RR{
|
||||||
|
&dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS}, Ns: "a.root-servers.net."},
|
||||||
|
&dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS}, Ns: "a.root-servers.net."},
|
||||||
|
}
|
||||||
|
names := extractNSRecords(rrs)
|
||||||
|
if len(names) != 1 {
|
||||||
|
t.Errorf("expected 1 deduped name, got %d", len(names))
|
||||||
|
}
|
||||||
|
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExtractIPsFromAnswer(t *testing.T) {
|
||||||
|
t.Run("empty", func(t *testing.T) {
|
||||||
|
ips := extractIPsFromAnswer(nil, dns.TypeA)
|
||||||
|
if len(ips) != 0 {
|
||||||
|
t.Errorf("expected 0 IPs, got %d", len(ips))
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("A records", func(t *testing.T) {
|
||||||
|
rrs := []dns.RR{
|
||||||
|
&dns.A{Hdr: dns.RR_Header{Rrtype: dns.TypeA}, A: net.ParseIP("198.41.0.4")},
|
||||||
|
&dns.A{Hdr: dns.RR_Header{Rrtype: dns.TypeA}, A: net.ParseIP("199.9.14.201")},
|
||||||
|
}
|
||||||
|
ips := extractIPsFromAnswer(rrs, dns.TypeA)
|
||||||
|
if len(ips) != 2 {
|
||||||
|
t.Fatalf("expected 2 IPs, got %d", len(ips))
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("AAAA records", func(t *testing.T) {
|
||||||
|
rrs := []dns.RR{
|
||||||
|
&dns.AAAA{Hdr: dns.RR_Header{Rrtype: dns.TypeAAAA}, AAAA: net.ParseIP("2001:503:ba3e::2:30")},
|
||||||
|
}
|
||||||
|
ips := extractIPsFromAnswer(rrs, dns.TypeAAAA)
|
||||||
|
if len(ips) != 1 {
|
||||||
|
t.Fatalf("expected 1 IP, got %d", len(ips))
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("filter by type", func(t *testing.T) {
|
||||||
|
rrs := []dns.RR{
|
||||||
|
&dns.A{Hdr: dns.RR_Header{Rrtype: dns.TypeA}, A: net.ParseIP("198.41.0.4")},
|
||||||
|
&dns.AAAA{Hdr: dns.RR_Header{Rrtype: dns.TypeAAAA}, AAAA: net.ParseIP("2001:503:ba3e::2:30")},
|
||||||
|
}
|
||||||
|
ips := extractIPsFromAnswer(rrs, dns.TypeA)
|
||||||
|
if len(ips) != 1 {
|
||||||
|
t.Fatalf("expected 1 A IP, got %d", len(ips))
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDiscoverRootsOverride(t *testing.T) {
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
cfg := &RootDiscoveryConfig{
|
||||||
|
Server: "a.root-servers.net",
|
||||||
|
IncludeAAAA: false,
|
||||||
|
}
|
||||||
|
|
||||||
|
servers, err := DiscoverRoots(ctx, cfg)
|
||||||
|
if err != nil {
|
||||||
|
t.Logf("skipping (no resolver available): %v", err)
|
||||||
|
t.Skip()
|
||||||
|
}
|
||||||
|
if len(servers) == 0 {
|
||||||
|
t.Fatal("expected at least one root server")
|
||||||
|
}
|
||||||
|
if len(servers[0].IPv4) == 0 {
|
||||||
|
t.Error("expected IPv4 addresses for a.root-servers.net")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDiscoverRootsSingle(t *testing.T) {
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
servers, err := DiscoverRoots(ctx, nil)
|
||||||
|
if err != nil {
|
||||||
|
t.Logf("skipping (no resolver available): %v", err)
|
||||||
|
t.Skip()
|
||||||
|
}
|
||||||
|
if len(servers) == 0 {
|
||||||
|
t.Fatal("expected at least one root server")
|
||||||
|
}
|
||||||
|
if servers[0].Name == "" {
|
||||||
|
t.Error("root server name should not be empty")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDiscoverRootsNilConfig(t *testing.T) {
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
servers, err := DiscoverRoots(ctx, nil)
|
||||||
|
if err != nil {
|
||||||
|
t.Logf("skipping (no resolver available): %v", err)
|
||||||
|
t.Skip()
|
||||||
|
}
|
||||||
|
if len(servers) == 0 {
|
||||||
|
t.Fatal("expected at least one root server with nil config")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildNSResponse(t *testing.T) {
|
||||||
|
msg := new(dns.Msg)
|
||||||
|
msg.SetReply(new(dns.Msg))
|
||||||
|
msg.Answer = append(msg.Answer,
|
||||||
|
&dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS, Class: dns.ClassINET}, Ns: "a.root-servers.net."},
|
||||||
|
&dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS, Class: dns.ClassINET}, Ns: "b.root-servers.net."},
|
||||||
|
&dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS, Class: dns.ClassINET}, Ns: "c.root-servers.net."},
|
||||||
|
)
|
||||||
|
|
||||||
|
names := extractNSRecords(msg.Answer)
|
||||||
|
if len(names) != 3 {
|
||||||
|
t.Fatalf("expected 3 NS records, got %d", len(names))
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, name := range names {
|
||||||
|
if len(name) == 0 || name[len(name)-1] != '.' {
|
||||||
|
t.Errorf("expected FQDN, got %q", name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user