feat: DNS host selection - fix code review issues (HAN-400) (#19)
CI / test (push) Failing after 1m58s

This commit was merged in pull request #19.
This commit is contained in:
2026-06-08 03:25:37 +00:00
parent 93959c9f11
commit d6063eb636
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")
rootAAAA := flag.Bool("root-aaaa", cfg.RootAAAA, "Include IPv6 root addresses")
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)")
allowTCP := flag.Bool("allow-tcp", cfg.AllowTCP, "TCP fallback on truncation")
alwaysTCP := flag.Bool("always-tcp", cfg.AlwaysTCP, "Always use TCP")
@@ -71,6 +72,7 @@ func main() {
cfg.AllRootServers = *allRootServers
cfg.RootAAAA = *rootAAAA
cfg.FollowAAAA = *followAAAA
cfg.DNSUpstream = *dnsUpstream
cfg.UDPSize = *udpSize
cfg.AllowTCP = *allowTCP
cfg.AlwaysTCP = *alwaysTCP
@@ -157,6 +159,7 @@ func main() {
IncludeAAAA: cfg.RootAAAA,
Server: rootServerAddr,
AllRoots: cfg.AllRootServers,
Resolver: cfg.DNSUpstream,
}
traverserConfig := &traverse.TraverserConfig{
@@ -179,6 +182,9 @@ func main() {
fmt.Fprintf(os.Stderr, " Allow TCP: %v\n", cfg.AllowTCP)
fmt.Fprintf(os.Stderr, " Always TCP: %v\n", cfg.AlwaysTCP)
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 {
+11
View File
@@ -28,6 +28,9 @@ type Config struct {
AllRootServers bool
RootAAAA bool
FollowAAAA bool
// DNSUpstream is the upstream resolver used for root server discovery.
// Format: "host:port" (e.g. "8.8.8.8:53"). Empty means use the system resolver.
DNSUpstream string
UDPSize int
AllowTCP bool
AlwaysTCP bool
@@ -127,6 +130,12 @@ func (c *Config) Validate() error {
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
}
@@ -173,6 +182,7 @@ func PrintUsage() {
{"--all-root-servers", "Use all 13 root servers"},
{"--root-aaaa", "Include IPv6 root addresses"},
{"--follow-aaaa", "Only follow AAAA for referrals"},
{"--dns-upstream", "Upstream resolver for root discovery (default: system)"},
},
"Transport Options": {
{"--udp-size", "EDNS0 buffer size (default 2048)"},
@@ -220,6 +230,7 @@ func DefaultConfig() *Config {
AllRootServers: false,
RootAAAA: false,
FollowAAAA: false,
DNSUpstream: "",
UDPSize: 2048,
AllowTCP: true,
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)
defer cancel()
servers, err := resolveRootServer(ctx, "a.root-servers.net.", false)
servers, err := resolveRootServer(ctx, systemResolver(), "a.root-servers.net.", false)
if err != nil {
t.Skipf("skipping (no local DNS): %v", err)
}
+61 -20
View File
@@ -26,7 +26,13 @@ func (rs *RootServer) AllIPs(includeAAAA bool) []net.IP {
}
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
}
@@ -43,20 +49,44 @@ func DiscoverRoots(ctx context.Context, cfg *RootDiscoveryConfig) ([]RootServer,
cfg = DefaultRootDiscoveryConfig()
}
resolver := resolverFromConfig(cfg)
if cfg.Server != "" {
return discoverRootOverride(ctx, cfg.Server, cfg.IncludeAAAA)
return discoverRootOverride(ctx, resolver, cfg.Server, cfg.IncludeAAAA)
}
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) {
resolver := systemResolver()
// 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)
@@ -70,16 +100,14 @@ func discoverRootOverride(ctx context.Context, server string, includeAAAA bool)
normalized := normalizeServerName(server)
for _, name := range nsSet {
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) {
resolver := systemResolver()
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)
@@ -95,12 +123,10 @@ func discoverSingleRoot(ctx context.Context, includeAAAA bool) ([]RootServer, er
}
pick := nsSet[0]
return resolveRootServer(ctx, pick, includeAAAA)
return resolveRootServer(ctx, resolver, pick, includeAAAA)
}
func discoverAllRoots(ctx context.Context, includeAAAA bool) ([]RootServer, error) {
resolver := systemResolver()
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)
@@ -117,7 +143,7 @@ func discoverAllRoots(ctx context.Context, includeAAAA bool) ([]RootServer, erro
var servers []RootServer
for _, name := range nsSet {
resolved, err := resolveRootServer(ctx, name, includeAAAA)
resolved, err := resolveRootServer(ctx, resolver, name, includeAAAA)
if err != nil {
servers = append(servers, RootServer{Name: name})
continue
@@ -132,9 +158,7 @@ func discoverAllRoots(ctx context.Context, includeAAAA bool) ([]RootServer, erro
return servers, nil
}
func resolveRootServer(ctx context.Context, name string, includeAAAA bool) ([]RootServer, error) {
resolver := systemResolver()
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 {
@@ -156,15 +180,32 @@ func resolveRootServer(ctx context.Context, name string, includeAAAA bool) ([]Ro
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,
WriteTimeout: 5,
ReadTimeout: 5 * time.Second,
WriteTimeout: 5 * time.Second,
}
if deadline, ok := ctx.Deadline(); ok {
c.ReadTimeout = time.Until(deadline)