feat: DNS host selection - fix code review issues (HAN-400) (#19)
CI / test (push) Failing after 1m58s
CI / test (push) Failing after 1m58s
This commit was merged in pull request #19.
This commit is contained in:
@@ -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 {
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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")},
|
||||||
|
},
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -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
@@ -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)
|
||||||
|
|||||||
Reference in New Issue
Block a user