Refactor code structure for improved readability and maintainability
This commit is contained in:
@@ -0,0 +1,255 @@
|
||||
package resolver
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"dns-helper/platform"
|
||||
)
|
||||
|
||||
// Resolve performs DNS resolution for hostname using the specified mode and config.
|
||||
// discoverer is used to obtain local network information (DNS servers, gateway).
|
||||
func Resolve(hostname string, mode ServerMode, config QueryConfig, discoverer platform.NetworkDiscoverer) ([]string, error) {
|
||||
var w io.Writer = io.Discard
|
||||
if config.Verbose {
|
||||
w = os.Stderr
|
||||
}
|
||||
|
||||
var ips []string
|
||||
var err error
|
||||
|
||||
switch mode.Mode {
|
||||
case "default", "":
|
||||
ips, err = resolveDefaultEntry(hostname, config, discoverer, w)
|
||||
case "local":
|
||||
ips, err = resolveLocal(hostname, config, discoverer)
|
||||
case "gateway":
|
||||
ips, err = resolveGateway(hostname, config, discoverer)
|
||||
case "explicit":
|
||||
ips, err = resolveExplicit(hostname, mode, config)
|
||||
default:
|
||||
return nil, fmt.Errorf("unknown server mode %q", mode.Mode)
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(ips) > 0 {
|
||||
fmt.Fprintf(w, "[dns] Result: %s\n", strings.Join(ips, ", "))
|
||||
}
|
||||
return ips, nil
|
||||
}
|
||||
|
||||
// resolveDefaultEntry is the public entry point for default mode.
|
||||
// It discovers network info once, builds the resolver pool, and delegates
|
||||
// to resolveDefault with depth=0.
|
||||
func resolveDefaultEntry(hostname string, config QueryConfig, discoverer platform.NetworkDiscoverer, w io.Writer) ([]string, error) {
|
||||
info, _ := discoverer.Discover() // FR-009: ignore discovery error, fall back to bootstrap
|
||||
pool := BuildResolverPool(ServerMode{Mode: "default"}, info)
|
||||
fmt.Fprintf(w, "[dns] Resolver pool: [%s]\n", strings.Join(pool, ", "))
|
||||
return resolveDefault(hostname, pool, info.DNSServers, config, 0, w)
|
||||
}
|
||||
|
||||
// resolveDefault runs the full smart-default resolution pipeline:
|
||||
// 1. Stage 1: Parallel NS fan-out across all resolver × label-level combinations.
|
||||
// 2. Stage 2: Non-recursive A query to the most-specific authoritative NS.
|
||||
// 2.5. Split-horizon cross-check (runs concurrently with Stage 2).
|
||||
// 3. Stage 3: Parallel A fallback if no NS records found.
|
||||
//
|
||||
// depth tracks CNAME chain hops; exceeding maxCNAMEDepth returns an error.
|
||||
func resolveDefault(hostname string, pool []string, localResolvers []string, config QueryConfig, depth int, w io.Writer) ([]string, error) {
|
||||
if depth > maxCNAMEDepth {
|
||||
return nil, fmt.Errorf(
|
||||
"CNAME chain depth exceeded for %s (max %d hops): probable CNAME loop or misconfigured zone",
|
||||
hostname, maxCNAMEDepth)
|
||||
}
|
||||
|
||||
labels := ExtractLabelLevels(hostname)
|
||||
|
||||
ctx := context.Background()
|
||||
timeout := config.Timeout
|
||||
if timeout <= 0 {
|
||||
timeout = 3 * time.Second
|
||||
}
|
||||
|
||||
// Stage 1: Parallel NS fan-out.
|
||||
fmt.Fprintf(w, "[dns] Stage 1: NS fan-out for %s (%d levels × %d resolvers = %d queries)\n",
|
||||
hostname, len(labels), len(pool), len(labels)*len(pool))
|
||||
nsResults := ParallelNSFanOut(ctx, w, pool, labels, timeout)
|
||||
authority := SelectAuthoritativeNS(nsResults)
|
||||
|
||||
if authority != nil {
|
||||
fmt.Fprintf(w, "[dns] Selected authority: %s → %s\n",
|
||||
authority.Zone, strings.Join(authority.Nameservers, ", "))
|
||||
return resolveAuthoritative(ctx, authority, hostname, pool, localResolvers, config, depth, w)
|
||||
}
|
||||
|
||||
// Stage 3: No NS records found — fall back to parallel A queries.
|
||||
fmt.Fprintf(w, "[dns] Stage 1: No NS records found at any level\n")
|
||||
ips, err := ParallelAFallback(ctx, w, pool, hostname, timeout)
|
||||
if err != nil {
|
||||
return nil, addPrivateTLDHint(hostname, err)
|
||||
}
|
||||
return ips, nil
|
||||
}
|
||||
|
||||
// resolveAuthoritative handles Stage 2 and Stage 2.5, plus CNAME restart.
|
||||
//
|
||||
// Stage 2.5 (split-horizon cross-check) is fired as a goroutine concurrently
|
||||
// with Stage 2 so it adds zero wall-clock latency to the happy path.
|
||||
func resolveAuthoritative(ctx context.Context, authority *AuthoritativeNS, hostname string, pool []string, localResolvers []string, config QueryConfig, depth int, w io.Writer) ([]string, error) {
|
||||
timeout := config.Timeout
|
||||
if timeout <= 0 {
|
||||
timeout = 3 * time.Second
|
||||
}
|
||||
|
||||
// Stage 2.5: fire local cross-check concurrently (FR-033, FR-037, FR-038).
|
||||
type shOut struct {
|
||||
localIPs []string
|
||||
localSrc string
|
||||
}
|
||||
var shCh chan shOut
|
||||
if len(localResolvers) > 0 {
|
||||
shCh = make(chan shOut, 1)
|
||||
go func() {
|
||||
localIPs, localSrc := queryLocalResolvers(ctx, localResolvers, hostname, timeout)
|
||||
shCh <- shOut{localIPs: localIPs, localSrc: localSrc}
|
||||
}()
|
||||
}
|
||||
|
||||
// Stage 2: Query authoritative NS with RD=false.
|
||||
ips, cnameTarget, nsHostnameUsed, err := QueryAuthoritative(ctx, w, authority, hostname, pool, timeout)
|
||||
if err != nil {
|
||||
return nil, addPrivateTLDHint(hostname, err)
|
||||
}
|
||||
|
||||
if cnameTarget != "" {
|
||||
// CNAME: restart full resolution for the target, incrementing depth.
|
||||
return resolveDefault(cnameTarget, pool, localResolvers, config, depth+1, w)
|
||||
}
|
||||
|
||||
// Stage 2.5: Collect cross-check result and compare IP sets.
|
||||
if shCh != nil {
|
||||
sh := <-shCh
|
||||
if sh.localIPs != nil {
|
||||
fmt.Fprintf(w, "[dns] Stage 2.5: Split-horizon cross-check against local resolvers [%s]\n",
|
||||
strings.Join(localResolvers, ", "))
|
||||
if !ipSetsEqual(ips, sh.localIPs) {
|
||||
fmt.Fprintf(w, "[dns] %s → %s (differs from authoritative %s)\n",
|
||||
sh.localSrc, strings.Join(sh.localIPs, ", "), strings.Join(ips, ", "))
|
||||
fmt.Fprintf(w, "[dns] CONFLICT: authoritative and local resolvers disagree\n")
|
||||
return nil, fmt.Errorf(
|
||||
"conflicting DNS answers for %s\n"+
|
||||
" Authoritative (%s): %s\n"+
|
||||
" Local resolver (%s): %s\n"+
|
||||
" The hostname resolves to different IPs depending on the DNS source.\n"+
|
||||
" Use -server local to trust your internal DNS, or -server <ip> to choose explicitly.",
|
||||
hostname,
|
||||
nsHostnameUsed, strings.Join(ips, ", "),
|
||||
sh.localSrc, strings.Join(sh.localIPs, ", "))
|
||||
}
|
||||
fmt.Fprintf(w, "[dns] %s → %s (matches authoritative)\n",
|
||||
sh.localSrc, strings.Join(sh.localIPs, ", "))
|
||||
fmt.Fprintf(w, "[dns] No conflict detected\n")
|
||||
}
|
||||
}
|
||||
|
||||
return ips, nil
|
||||
}
|
||||
|
||||
// queryLocalResolvers tries each localResolver in order and returns the first
|
||||
// successful A record result. Returns nil, "" if all fail or return NXDOMAIN.
|
||||
func queryLocalResolvers(ctx context.Context, localResolvers []string, hostname string, timeout time.Duration) ([]string, string) {
|
||||
for _, lr := range localResolvers {
|
||||
queryCtx, cancel := context.WithTimeout(ctx, timeout)
|
||||
ips, err := queryA(queryCtx, lr, hostname, timeout)
|
||||
cancel()
|
||||
if err == nil && len(ips) > 0 {
|
||||
return ips, lr
|
||||
}
|
||||
}
|
||||
return nil, ""
|
||||
}
|
||||
|
||||
// resolveLocal queries only the locally configured DNS servers in priority order.
|
||||
// No public resolver fallback is used (FR-026).
|
||||
func resolveLocal(hostname string, config QueryConfig, discoverer platform.NetworkDiscoverer) ([]string, error) {
|
||||
info, err := discoverer.Discover()
|
||||
if err != nil || len(info.DNSServers) == 0 {
|
||||
return nil, fmt.Errorf("failed to resolve %s using local resolvers: no local DNS servers found", hostname)
|
||||
}
|
||||
|
||||
timeout := config.Timeout
|
||||
if timeout <= 0 {
|
||||
timeout = 3 * time.Second
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
for _, server := range info.DNSServers {
|
||||
ips, queryErr := queryA(ctx, server, hostname, timeout)
|
||||
if queryErr == nil && len(ips) > 0 {
|
||||
return ips, nil
|
||||
}
|
||||
}
|
||||
|
||||
return nil, fmt.Errorf(
|
||||
"failed to resolve %s using local resolvers: all local DNS servers are unreachable\n Local resolvers tried: %s",
|
||||
hostname, strings.Join(info.DNSServers, ", "))
|
||||
}
|
||||
|
||||
// resolveGateway queries the default gateway as a DNS server.
|
||||
// No other resolvers are tried (FR-027).
|
||||
func resolveGateway(hostname string, config QueryConfig, discoverer platform.NetworkDiscoverer) ([]string, error) {
|
||||
info, err := discoverer.Discover()
|
||||
if err != nil || info.Gateway == "" {
|
||||
return nil, fmt.Errorf("failed to resolve %s using gateway: no default gateway found", hostname)
|
||||
}
|
||||
|
||||
timeout := config.Timeout
|
||||
if timeout <= 0 {
|
||||
timeout = 3 * time.Second
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
ips, queryErr := queryA(ctx, info.Gateway, hostname, timeout)
|
||||
if queryErr != nil {
|
||||
return nil, fmt.Errorf(
|
||||
"failed to resolve %s using gateway: gateway %s did not respond to DNS query",
|
||||
hostname, info.Gateway)
|
||||
}
|
||||
return ips, nil
|
||||
}
|
||||
|
||||
// resolveExplicit queries the server address provided directly by the user.
|
||||
// Uses the per-query timeout from config to respect the -timeout flag (FR-029).
|
||||
func resolveExplicit(hostname string, mode ServerMode, config QueryConfig) ([]string, error) {
|
||||
timeout := config.Timeout
|
||||
if timeout <= 0 {
|
||||
timeout = 3 * time.Second
|
||||
}
|
||||
ctx := context.Background()
|
||||
ips, err := queryA(ctx, mode.ExplicitAddr, hostname, timeout)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to resolve %s via %s: %w", hostname, mode.ExplicitAddr, err)
|
||||
}
|
||||
return ips, nil
|
||||
}
|
||||
|
||||
// addPrivateTLDHint wraps err with a user-friendly hint when the hostname uses
|
||||
// a private TLD (FR-025).
|
||||
func addPrivateTLDHint(hostname string, origErr error) error {
|
||||
parts := strings.Split(strings.TrimSuffix(hostname, "."), ".")
|
||||
if len(parts) > 0 {
|
||||
tld := strings.ToLower(parts[len(parts)-1])
|
||||
if PrivateTLDs[tld] {
|
||||
return fmt.Errorf(
|
||||
"%w\n Hint: the hostname uses a private TLD (.%s). Try specifying an internal DNS server:\n dns-helper add -host %s -server <internal-dns-ip>",
|
||||
origErr, tld, hostname)
|
||||
}
|
||||
}
|
||||
return origErr
|
||||
}
|
||||
Reference in New Issue
Block a user