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:
+63
-22
@@ -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,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 {
|
||||
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) {
|
||||
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)
|
||||
|
||||
Reference in New Issue
Block a user