feat: DNS host selection - fix code review issues (HAN-400) #19

Merged
multica-agent merged 2 commits from agent/go-expert-developer/96b6d6e2 into main 2026-06-08 03:25:38 +00:00
6 changed files with 321 additions and 32 deletions
+6
View File
@@ -21,6 +21,7 @@ func main() {
allRootServers := flag.Bool("all-root-servers", cfg.AllRootServers, "Use all 13 root servers") allRootServers := flag.Bool("all-root-servers", cfg.AllRootServers, "Use all 13 root servers")
rootAAAA := flag.Bool("root-aaaa", cfg.RootAAAA, "Include IPv6 root addresses") rootAAAA := flag.Bool("root-aaaa", cfg.RootAAAA, "Include IPv6 root addresses")
followAAAA := flag.Bool("follow-aaaa", cfg.FollowAAAA, "Only follow AAAA for referrals") followAAAA := flag.Bool("follow-aaaa", cfg.FollowAAAA, "Only follow AAAA for referrals")
dnsUpstream := flag.String("dns-upstream", cfg.DNSUpstream, "Upstream resolver for root discovery (e.g. 8.8.8.8:53, default: system)")
udpSize := flag.Int("udp-size", cfg.UDPSize, "EDNS0 buffer size (512-4096)") udpSize := flag.Int("udp-size", cfg.UDPSize, "EDNS0 buffer size (512-4096)")
allowTCP := flag.Bool("allow-tcp", cfg.AllowTCP, "TCP fallback on truncation") allowTCP := flag.Bool("allow-tcp", cfg.AllowTCP, "TCP fallback on truncation")
alwaysTCP := flag.Bool("always-tcp", cfg.AlwaysTCP, "Always use TCP") alwaysTCP := flag.Bool("always-tcp", cfg.AlwaysTCP, "Always use TCP")
@@ -71,6 +72,7 @@ func main() {
cfg.AllRootServers = *allRootServers cfg.AllRootServers = *allRootServers
cfg.RootAAAA = *rootAAAA cfg.RootAAAA = *rootAAAA
cfg.FollowAAAA = *followAAAA cfg.FollowAAAA = *followAAAA
cfg.DNSUpstream = *dnsUpstream
cfg.UDPSize = *udpSize cfg.UDPSize = *udpSize
cfg.AllowTCP = *allowTCP cfg.AllowTCP = *allowTCP
cfg.AlwaysTCP = *alwaysTCP cfg.AlwaysTCP = *alwaysTCP
@@ -157,6 +159,7 @@ func main() {
IncludeAAAA: cfg.RootAAAA, IncludeAAAA: cfg.RootAAAA,
Server: rootServerAddr, Server: rootServerAddr,
AllRoots: cfg.AllRootServers, AllRoots: cfg.AllRootServers,
Resolver: cfg.DNSUpstream,
} }
traverserConfig := &traverse.TraverserConfig{ traverserConfig := &traverse.TraverserConfig{
@@ -179,6 +182,9 @@ func main() {
fmt.Fprintf(os.Stderr, " Allow TCP: %v\n", cfg.AllowTCP) fmt.Fprintf(os.Stderr, " Allow TCP: %v\n", cfg.AllowTCP)
fmt.Fprintf(os.Stderr, " Always TCP: %v\n", cfg.AlwaysTCP) fmt.Fprintf(os.Stderr, " Always TCP: %v\n", cfg.AlwaysTCP)
fmt.Fprintf(os.Stderr, " Fast Mode: %v\n", cfg.Fast) fmt.Fprintf(os.Stderr, " Fast Mode: %v\n", cfg.Fast)
if cfg.DNSUpstream != "" {
fmt.Fprintf(os.Stderr, " DNS Upstream: %s\n", cfg.DNSUpstream)
}
} }
if !cfg.Quiet && !*jsonOutput { if !cfg.Quiet && !*jsonOutput {
+20 -9
View File
@@ -28,15 +28,18 @@ type Config struct {
AllRootServers bool AllRootServers bool
RootAAAA bool RootAAAA bool
FollowAAAA bool FollowAAAA bool
UDPSize int // DNSUpstream is the upstream resolver used for root server discovery.
AllowTCP bool // Format: "host:port" (e.g. "8.8.8.8:53"). Empty means use the system resolver.
AlwaysTCP bool DNSUpstream string
MaxDepth int UDPSize int
Retries int AllowTCP bool
Fast bool AlwaysTCP bool
Verbose bool MaxDepth int
Debug int Retries int
Quiet bool Fast bool
Verbose bool
Debug int
Quiet bool
ShowProgress bool ShowProgress bool
ShowResolves bool ShowResolves bool
@@ -127,6 +130,12 @@ func (c *Config) Validate() error {
return ErrAlwaysTCPRequiresTCP return ErrAlwaysTCPRequiresTCP
} }
if c.DNSUpstream != "" {
if _, _, err := net.SplitHostPort(c.DNSUpstream); err != nil {
return fmt.Errorf("--dns-upstream %q is not a valid host:port address", c.DNSUpstream)
}
}
return nil return nil
} }
@@ -173,6 +182,7 @@ func PrintUsage() {
{"--all-root-servers", "Use all 13 root servers"}, {"--all-root-servers", "Use all 13 root servers"},
{"--root-aaaa", "Include IPv6 root addresses"}, {"--root-aaaa", "Include IPv6 root addresses"},
{"--follow-aaaa", "Only follow AAAA for referrals"}, {"--follow-aaaa", "Only follow AAAA for referrals"},
{"--dns-upstream", "Upstream resolver for root discovery (default: system)"},
}, },
"Transport Options": { "Transport Options": {
{"--udp-size", "EDNS0 buffer size (default 2048)"}, {"--udp-size", "EDNS0 buffer size (default 2048)"},
@@ -220,6 +230,7 @@ func DefaultConfig() *Config {
AllRootServers: false, AllRootServers: false,
RootAAAA: false, RootAAAA: false,
FollowAAAA: false, FollowAAAA: false,
DNSUpstream: "",
UDPSize: 2048, UDPSize: 2048,
AllowTCP: true, AllowTCP: true,
AlwaysTCP: false, AlwaysTCP: false,
+74
View File
@@ -0,0 +1,74 @@
package dns
import "net"
// RootHints contains the 13 IANA root name servers with their well-known
// IPv4 and IPv6 addresses as published at https://www.iana.org/domains/root/servers.
// These addresses change very rarely and are safe to embed as application constants.
var RootHints = []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")},
},
{
Name: "b.root-servers.net.",
IPv4: []net.IP{net.ParseIP("170.247.170.2")},
IPv6: []net.IP{net.ParseIP("2801:1b8:10::b")},
},
{
Name: "c.root-servers.net.",
IPv4: []net.IP{net.ParseIP("192.33.4.12")},
IPv6: []net.IP{net.ParseIP("2001:500:2::c")},
},
{
Name: "d.root-servers.net.",
IPv4: []net.IP{net.ParseIP("199.7.91.13")},
IPv6: []net.IP{net.ParseIP("2001:500:2d::d")},
},
{
Name: "e.root-servers.net.",
IPv4: []net.IP{net.ParseIP("192.203.230.10")},
IPv6: []net.IP{net.ParseIP("2001:500:a8::e")},
},
{
Name: "f.root-servers.net.",
IPv4: []net.IP{net.ParseIP("192.5.5.241")},
IPv6: []net.IP{net.ParseIP("2001:500:2f::f")},
},
{
Name: "g.root-servers.net.",
IPv4: []net.IP{net.ParseIP("192.112.36.4")},
IPv6: []net.IP{net.ParseIP("2001:500:12::d0d")},
},
{
Name: "h.root-servers.net.",
IPv4: []net.IP{net.ParseIP("198.97.190.53")},
IPv6: []net.IP{net.ParseIP("2001:500:1::53")},
},
{
Name: "i.root-servers.net.",
IPv4: []net.IP{net.ParseIP("192.36.148.17")},
IPv6: []net.IP{net.ParseIP("2001:7fe::53")},
},
{
Name: "j.root-servers.net.",
IPv4: []net.IP{net.ParseIP("192.58.128.30")},
IPv6: []net.IP{net.ParseIP("2001:503:c27::2:30")},
},
{
Name: "k.root-servers.net.",
IPv4: []net.IP{net.ParseIP("193.0.14.129")},
IPv6: []net.IP{net.ParseIP("2001:7fd::1")},
},
{
Name: "l.root-servers.net.",
IPv4: []net.IP{net.ParseIP("199.7.83.42")},
IPv6: []net.IP{net.ParseIP("2001:500:9f::42")},
},
{
Name: "m.root-servers.net.",
IPv4: []net.IP{net.ParseIP("202.12.27.33")},
IPv6: []net.IP{net.ParseIP("2001:dc3::35")},
},
}
+157
View File
@@ -0,0 +1,157 @@
package dns
import (
"net"
"testing"
)
func TestRootHintsCount(t *testing.T) {
if len(RootHints) != 13 {
t.Errorf("expected 13 root hints, got %d", len(RootHints))
}
}
func TestRootHintsNames(t *testing.T) {
wantNames := []string{
"a.root-servers.net.",
"b.root-servers.net.",
"c.root-servers.net.",
"d.root-servers.net.",
"e.root-servers.net.",
"f.root-servers.net.",
"g.root-servers.net.",
"h.root-servers.net.",
"i.root-servers.net.",
"j.root-servers.net.",
"k.root-servers.net.",
"l.root-servers.net.",
"m.root-servers.net.",
}
for i, rs := range RootHints {
if rs.Name != wantNames[i] {
t.Errorf("RootHints[%d].Name = %q, want %q", i, rs.Name, wantNames[i])
}
}
}
func TestRootHintsHaveIPv4(t *testing.T) {
for _, rs := range RootHints {
if len(rs.IPv4) == 0 {
t.Errorf("root server %q has no IPv4 address", rs.Name)
}
for _, ip := range rs.IPv4 {
if ip.To4() == nil {
t.Errorf("root server %q: expected IPv4, got %v", rs.Name, ip)
}
}
}
}
func TestRootHintsHaveIPv6(t *testing.T) {
for _, rs := range RootHints {
if len(rs.IPv6) == 0 {
t.Errorf("root server %q has no IPv6 address", rs.Name)
}
for _, ip := range rs.IPv6 {
if ip.To4() != nil {
t.Errorf("root server %q: expected IPv6, got IPv4-mappable %v", rs.Name, ip)
}
}
}
}
func TestRootHintsAllIPsIPv4Only(t *testing.T) {
for _, rs := range RootHints {
ips := rs.AllIPs(false)
if len(ips) != len(rs.IPv4) {
t.Errorf("root server %q: AllIPs(false) = %d, want %d", rs.Name, len(ips), len(rs.IPv4))
}
}
}
func TestRootHintsAllIPsBoth(t *testing.T) {
for _, rs := range RootHints {
ips := rs.AllIPs(true)
want := len(rs.IPv4) + len(rs.IPv6)
if len(ips) != want {
t.Errorf("root server %q: AllIPs(true) = %d, want %d", rs.Name, len(ips), want)
}
}
}
func TestRootHintsNoParseFail(t *testing.T) {
// Ensure none of the IPs failed to parse (net.ParseIP returns nil on failure).
for _, rs := range RootHints {
for _, ip := range rs.IPv4 {
if ip == nil {
t.Errorf("root server %q has nil IPv4 (parse failed)", rs.Name)
}
}
for _, ip := range rs.IPv6 {
if ip == nil {
t.Errorf("root server %q has nil IPv6 (parse failed)", rs.Name)
}
}
}
}
func TestRootHintsKnownAddress(t *testing.T) {
// Spot-check a.root-servers.net. which has been stable for decades.
for _, rs := range RootHints {
if rs.Name == "a.root-servers.net." {
want := net.ParseIP("198.41.0.4")
if !rs.IPv4[0].Equal(want) {
t.Errorf("a.root-servers.net. IPv4 = %v, want %v", rs.IPv4[0], want)
}
return
}
}
t.Error("a.root-servers.net. not found in RootHints")
}
func TestResolverFromConfigExplicit(t *testing.T) {
cfg := &RootDiscoveryConfig{Resolver: "8.8.8.8:53"}
got := resolverFromConfig(cfg)
if got != "8.8.8.8:53" {
t.Errorf("resolverFromConfig = %q, want 8.8.8.8:53", got)
}
}
func TestResolverFromConfigEmpty(t *testing.T) {
cfg := &RootDiscoveryConfig{}
got := resolverFromConfig(cfg)
// Should return the system resolver; just check it's non-empty and contains a port.
if got == "" {
t.Error("resolverFromConfig with empty Resolver returned empty string")
}
}
func TestResolverFromConfigNil(t *testing.T) {
got := resolverFromConfig(nil)
if got == "" {
t.Error("resolverFromConfig(nil) returned empty string")
}
}
func TestSystemResolverNonEmpty(t *testing.T) {
got := systemResolver()
if got == "" {
t.Error("systemResolver() returned empty string")
}
// Must contain a colon (host:port format).
host, port, err := splitHostPort(got)
if err != nil {
t.Errorf("systemResolver() = %q: not a valid host:port: %v", got, err)
}
if host == "" {
t.Errorf("systemResolver() host is empty in %q", got)
}
if port == "" {
t.Errorf("systemResolver() port is empty in %q", got)
}
}
// splitHostPort is a thin wrapper around net.SplitHostPort for test use.
func splitHostPort(addr string) (host, port string, err error) {
return net.SplitHostPort(addr)
}
+1 -1
View File
@@ -217,7 +217,7 @@ func TestResolveRootServerDirect(t *testing.T) {
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel() defer cancel()
servers, err := resolveRootServer(ctx, "a.root-servers.net.", false) servers, err := resolveRootServer(ctx, systemResolver(), "a.root-servers.net.", false)
if err != nil { if err != nil {
t.Skipf("skipping (no local DNS): %v", err) t.Skipf("skipping (no local DNS): %v", err)
} }
+63 -22
View File
@@ -26,7 +26,13 @@ func (rs *RootServer) AllIPs(includeAAAA bool) []net.IP {
} }
type RootDiscoveryConfig struct { type RootDiscoveryConfig struct {
Server string // 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 AllRoots bool
IncludeAAAA bool IncludeAAAA bool
} }
@@ -43,20 +49,44 @@ func DiscoverRoots(ctx context.Context, cfg *RootDiscoveryConfig) ([]RootServer,
cfg = DefaultRootDiscoveryConfig() cfg = DefaultRootDiscoveryConfig()
} }
resolver := resolverFromConfig(cfg)
if cfg.Server != "" { if cfg.Server != "" {
return discoverRootOverride(ctx, cfg.Server, cfg.IncludeAAAA) return discoverRootOverride(ctx, resolver, cfg.Server, cfg.IncludeAAAA)
} }
if cfg.AllRoots { if cfg.AllRoots {
return discoverAllRoots(ctx, cfg.IncludeAAAA) servers, err := discoverAllRoots(ctx, resolver, cfg.IncludeAAAA)
if err != nil {
return filterHints(RootHints, cfg.IncludeAAAA), nil
}
return servers, nil
} }
return discoverSingleRoot(ctx, cfg.IncludeAAAA) 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
} }
func discoverRootOverride(ctx context.Context, server string, includeAAAA bool) ([]RootServer, error) { // filterHints returns a copy of hints with IPv6 addresses stripped when includeAAAA is false.
resolver := systemResolver() 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) nsMsg, err := queryResolver(ctx, resolver, ".", dns.TypeNS)
if err != nil { if err != nil {
return nil, fmt.Errorf("query root NS records: %w", err) return nil, fmt.Errorf("query root NS records: %w", err)
@@ -70,16 +100,14 @@ func discoverRootOverride(ctx context.Context, server string, includeAAAA bool)
normalized := normalizeServerName(server) normalized := normalizeServerName(server)
for _, name := range nsSet { for _, name := range nsSet {
if normalizeServerName(name) == normalized { if normalizeServerName(name) == normalized {
return resolveRootServer(ctx, name, includeAAAA) return resolveRootServer(ctx, resolver, name, includeAAAA)
} }
} }
return resolveRootServer(ctx, server, includeAAAA) return resolveRootServer(ctx, resolver, server, includeAAAA)
} }
func discoverSingleRoot(ctx context.Context, includeAAAA bool) ([]RootServer, error) { func discoverSingleRoot(ctx context.Context, resolver string, includeAAAA bool) ([]RootServer, error) {
resolver := systemResolver()
nsMsg, err := queryResolver(ctx, resolver, ".", dns.TypeNS) nsMsg, err := queryResolver(ctx, resolver, ".", dns.TypeNS)
if err != nil { if err != nil {
return nil, fmt.Errorf("query root NS records: %w", err) return nil, fmt.Errorf("query root NS records: %w", err)
@@ -95,12 +123,10 @@ func discoverSingleRoot(ctx context.Context, includeAAAA bool) ([]RootServer, er
} }
pick := nsSet[0] pick := nsSet[0]
return resolveRootServer(ctx, pick, includeAAAA) return resolveRootServer(ctx, resolver, pick, includeAAAA)
} }
func discoverAllRoots(ctx context.Context, includeAAAA bool) ([]RootServer, error) { func discoverAllRoots(ctx context.Context, resolver string, includeAAAA bool) ([]RootServer, error) {
resolver := systemResolver()
nsMsg, err := queryResolver(ctx, resolver, ".", dns.TypeNS) nsMsg, err := queryResolver(ctx, resolver, ".", dns.TypeNS)
if err != nil { if err != nil {
return nil, fmt.Errorf("query root NS records: %w", err) return nil, fmt.Errorf("query root NS records: %w", err)
@@ -117,7 +143,7 @@ func discoverAllRoots(ctx context.Context, includeAAAA bool) ([]RootServer, erro
var servers []RootServer var servers []RootServer
for _, name := range nsSet { for _, name := range nsSet {
resolved, err := resolveRootServer(ctx, name, includeAAAA) resolved, err := resolveRootServer(ctx, resolver, name, includeAAAA)
if err != nil { if err != nil {
servers = append(servers, RootServer{Name: name}) servers = append(servers, RootServer{Name: name})
continue continue
@@ -132,9 +158,7 @@ func discoverAllRoots(ctx context.Context, includeAAAA bool) ([]RootServer, erro
return servers, nil return servers, nil
} }
func resolveRootServer(ctx context.Context, name string, includeAAAA bool) ([]RootServer, error) { func resolveRootServer(ctx context.Context, resolver, name string, includeAAAA bool) ([]RootServer, error) {
resolver := systemResolver()
var ipv4 []net.IP var ipv4 []net.IP
aMsg, err := queryResolver(ctx, resolver, name, dns.TypeA) aMsg, err := queryResolver(ctx, resolver, name, dns.TypeA)
if err == nil { if err == nil {
@@ -156,15 +180,32 @@ func resolveRootServer(ctx context.Context, name string, includeAAAA bool) ([]Ro
return []RootServer{{Name: name, IPv4: ipv4, IPv6: ipv6}}, nil 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 { func systemResolver() string {
return "127.0.0.1:53" 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) { func queryResolver(ctx context.Context, resolverAddr, name string, qtype uint16) (*dns.Msg, error) {
c := &dns.Client{ c := &dns.Client{
Net: "udp", Net: "udp",
ReadTimeout: 5, ReadTimeout: 5 * time.Second,
WriteTimeout: 5, WriteTimeout: 5 * time.Second,
} }
if deadline, ok := ctx.Deadline(); ok { if deadline, ok := ctx.Deadline(); ok {
c.ReadTimeout = time.Until(deadline) c.ReadTimeout = time.Until(deadline)