diff --git a/cmd/exploredns/main.go b/cmd/exploredns/main.go index bd50c42..6e64d49 100644 --- a/cmd/exploredns/main.go +++ b/cmd/exploredns/main.go @@ -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 { diff --git a/internal/config/config.go b/internal/config/config.go index 4d7cf0e..7bb9ebc 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -28,15 +28,18 @@ type Config struct { AllRootServers bool RootAAAA bool FollowAAAA bool - UDPSize int - AllowTCP bool - AlwaysTCP bool - MaxDepth int - Retries int - Fast bool - Verbose bool - Debug int - Quiet 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 + MaxDepth int + Retries int + Fast bool + Verbose bool + Debug int + Quiet bool ShowProgress bool ShowResolves bool @@ -173,6 +176,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 +224,7 @@ func DefaultConfig() *Config { AllRootServers: false, RootAAAA: false, FollowAAAA: false, + DNSUpstream: "", UDPSize: 2048, AllowTCP: true, AlwaysTCP: false, diff --git a/internal/dns/hints.go b/internal/dns/hints.go new file mode 100644 index 0000000..1cf6ebc --- /dev/null +++ b/internal/dns/hints.go @@ -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")}, + }, +} diff --git a/internal/dns/hints_test.go b/internal/dns/hints_test.go new file mode 100644 index 0000000..307f87c --- /dev/null +++ b/internal/dns/hints_test.go @@ -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) +} diff --git a/internal/dns/real_exchange_test.go b/internal/dns/real_exchange_test.go index c111b26..ec43bf4 100644 --- a/internal/dns/real_exchange_test.go +++ b/internal/dns/real_exchange_test.go @@ -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) } diff --git a/internal/dns/roots.go b/internal/dns/roots.go index 811eafe..fac629b 100644 --- a/internal/dns/roots.go +++ b/internal/dns/roots.go @@ -26,7 +26,13 @@ func (rs *RootServer) AllIPs(includeAAAA bool) []net.IP { } 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 IncludeAAAA bool } @@ -43,20 +49,20 @@ 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) + return discoverAllRoots(ctx, resolver, cfg.IncludeAAAA) } - return discoverSingleRoot(ctx, cfg.IncludeAAAA) + return discoverSingleRoot(ctx, resolver, cfg.IncludeAAAA) } -func discoverRootOverride(ctx context.Context, server string, includeAAAA bool) ([]RootServer, error) { - resolver := systemResolver() - +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 +76,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 +99,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 +119,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 +134,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,8 +156,25 @@ 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. +// On Unix-like systems this reads /etc/resolv.conf. Falls back to 127.0.0.1:53 +// when the system configuration is unavailable or contains no servers. 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) {