package resolver import ( "context" "fmt" "io" "os" "strings" "time" "ekdns/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 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 ekdns add -host %s -server ", origErr, tld, hostname) } } return origErr }