package resolver import ( "context" "fmt" "io" "math/rand" "net" "sort" "strings" "time" "golang.org/x/net/dns/dnsmessage" ) // NSResult is the result of a single NS query sent during the parallel fan-out. type NSResult struct { LabelLevel string // Domain level queried (e.g., "example.com") Resolver string // Resolver address that was queried NSRecords []string // NS hostnames returned (nil if none) CNAMETarget string // CNAME target if the NS query returned a CNAME Err error // Transport or parse error (nil for NXDOMAIN) } // AuthoritativeNS holds the zone and authoritative nameservers selected from // a fan-out result set. type AuthoritativeNS struct { Zone string // Domain level that returned NS records (e.g., "example.com") Nameservers []string // Sorted, deduplicated NS hostnames } // ParallelNSFanOut queries NS records for every combination of (resolver, labelLevel) // in parallel. Returns one NSResult per combination — total len(resolvers)*len(labelLevels). // w receives per-level NS result summaries after all queries complete (pass io.Discard to suppress). func ParallelNSFanOut(ctx context.Context, w io.Writer, resolvers []string, labelLevels []string, timeout time.Duration) []NSResult { total := len(resolvers) * len(labelLevels) if total == 0 { return nil } ch := make(chan NSResult, total) for _, r := range resolvers { for _, level := range labelLevels { go func(resolver, label string) { ch <- queryNS(ctx, resolver, label, timeout) }(r, level) } } results := make([]NSResult, 0, total) for i := 0; i < total; i++ { results = append(results, <-ch) } // Verbose: log per-level NS result summary after all queries have returned. for _, lvl := range labelLevels { var nsSet map[string]bool var viaResolvers []string for _, r := range results { if r.LabelLevel != lvl { continue } if len(r.NSRecords) > 0 { if nsSet == nil { nsSet = make(map[string]bool) } for _, ns := range r.NSRecords { nsSet[ns] = true } viaResolvers = append(viaResolvers, r.Resolver) } } if len(nsSet) > 0 { nsList := make([]string, 0, len(nsSet)) for ns := range nsSet { nsList = append(nsList, ns) } sort.Strings(nsList) sort.Strings(viaResolvers) uniqueVia := viaResolvers[:0:0] for i, rv := range viaResolvers { if i == 0 || rv != viaResolvers[i-1] { uniqueVia = append(uniqueVia, rv) } } fmt.Fprintf(w, "[dns] %s NS: %s (via %s)\n", lvl, strings.Join(nsList, ", "), strings.Join(uniqueVia, ", ")) } else { fmt.Fprintf(w, "[dns] %s NS: (none)\n", lvl) } } return results } // queryNS sends a single NS query for label to resolver and returns an NSResult. // NXDOMAIN is not an error — it returns an NSResult with nil NSRecords and nil Err. func queryNS(ctx context.Context, resolver, label string, timeout time.Duration) NSResult { fqdn := label if !strings.HasSuffix(fqdn, ".") { fqdn += "." } name, err := dnsmessage.NewName(fqdn) if err != nil { return NSResult{LabelLevel: label, Resolver: resolver, Err: fmt.Errorf("invalid hostname %q: %w", label, err)} } id := uint16(rand.Uint32()) //nolint:gosec msg := dnsmessage.Message{ Header: dnsmessage.Header{ ID: id, RecursionDesired: true, }, Questions: []dnsmessage.Question{{ Name: name, Type: dnsmessage.TypeNS, Class: dnsmessage.ClassINET, }}, } packed, err := msg.Pack() if err != nil { return NSResult{LabelLevel: label, Resolver: resolver, Err: fmt.Errorf("packing NS query: %w", err)} } queryCtx, cancel := context.WithTimeout(ctx, timeout) defer cancel() resp, err := UDPQuery(queryCtx, resolver, packed, timeout) if err != nil { return NSResult{LabelLevel: label, Resolver: resolver, Err: err} } var parser dnsmessage.Parser respHeader, err := parser.Start(resp) if err != nil { return NSResult{LabelLevel: label, Resolver: resolver, Err: fmt.Errorf("parsing NS response: %w", err)} } if respHeader.ID != id { return NSResult{LabelLevel: label, Resolver: resolver, Err: fmt.Errorf("NS response ID mismatch (expected %d, got %d)", id, respHeader.ID)} } // NXDOMAIN: not an error, just no NS records at this level. if respHeader.RCode == dnsmessage.RCodeNameError { return NSResult{LabelLevel: label, Resolver: resolver} } if respHeader.RCode != dnsmessage.RCodeSuccess { return NSResult{LabelLevel: label, Resolver: resolver, Err: fmt.Errorf("NS query for %s returned %s", label, respHeader.RCode)} } if err := parser.SkipAllQuestions(); err != nil { return NSResult{LabelLevel: label, Resolver: resolver, Err: fmt.Errorf("skipping questions in NS response: %w", err)} } var nsRecords []string var cnameTarget string for { hdr, err := parser.AnswerHeader() if err == dnsmessage.ErrSectionDone { break } if err != nil { return NSResult{LabelLevel: label, Resolver: resolver, Err: fmt.Errorf("parsing NS answer header: %w", err)} } switch hdr.Type { case dnsmessage.TypeNS: nsRec, err := parser.NSResource() if err != nil { return NSResult{LabelLevel: label, Resolver: resolver, Err: fmt.Errorf("parsing NS record: %w", err)} } nsRecords = append(nsRecords, nsRec.NS.String()) case dnsmessage.TypeCNAME: cnameRec, err := parser.CNAMEResource() if err != nil { return NSResult{LabelLevel: label, Resolver: resolver, Err: fmt.Errorf("parsing CNAME in NS response: %w", err)} } // Preserve without stripping dot — caller may need it. cnameTarget = strings.TrimSuffix(cnameRec.CNAME.String(), ".") default: if err := parser.SkipAnswer(); err != nil { return NSResult{LabelLevel: label, Resolver: resolver, Err: fmt.Errorf("skipping NS answer: %w", err)} } } } return NSResult{ LabelLevel: label, Resolver: resolver, NSRecords: nsRecords, CNAMETarget: cnameTarget, } } // SelectAuthoritativeNS picks the most-specific NS zone from the fan-out results. // NS records across resolvers for the same zone are merged and deduplicated. // Returns nil if no NS records were found at any level. func SelectAuthoritativeNS(results []NSResult) *AuthoritativeNS { type levelInfo struct { nsSet map[string]bool } levels := make(map[string]*levelInfo) for _, r := range results { if len(r.NSRecords) == 0 { continue } if _, ok := levels[r.LabelLevel]; !ok { levels[r.LabelLevel] = &levelInfo{nsSet: make(map[string]bool)} } for _, ns := range r.NSRecords { levels[r.LabelLevel].nsSet[ns] = true } } if len(levels) == 0 { return nil } // Select the most specific zone: most label segments wins. // Tie-break lexicographically (deterministic test output). bestLabel := "" bestCount := 0 for level := range levels { count := strings.Count(level, ".") + 1 if count > bestCount || (count == bestCount && level > bestLabel) { bestCount = count bestLabel = level } } nsSet := levels[bestLabel].nsSet nameservers := make([]string, 0, len(nsSet)) for ns := range nsSet { nameservers = append(nameservers, ns) } sort.Strings(nameservers) return &AuthoritativeNS{Zone: bestLabel, Nameservers: nameservers} } // QueryAuthoritative sends a non-recursive A query to each authoritative // nameserver in turn, returning the first successful result. // // resolvers is the full resolver pool used to resolve NS hostnames to IPs. // w receives Stage 2 diagnostic lines (pass io.Discard to suppress). // Returns: ips, cnameTarget, nsHostnameUsed, error. // If a CNAME is returned, ips is nil and cnameTarget is populated — caller restarts. func QueryAuthoritative(ctx context.Context, w io.Writer, ns *AuthoritativeNS, hostname string, resolvers []string, timeout time.Duration) (ips []string, cnameTarget string, nsHostnameUsed string, err error) { fqdn := hostname if !strings.HasSuffix(fqdn, ".") { fqdn += "." } qname, nameErr := dnsmessage.NewName(fqdn) if nameErr != nil { return nil, "", "", fmt.Errorf("invalid hostname %q: %w", hostname, nameErr) } for _, nsHostname := range ns.Nameservers { // Step 1: Resolve NS hostname to an IP using the resolver pool. nsIP, resolveErr := resolveNSHostname(ctx, nsHostname, resolvers, timeout) if resolveErr != nil { continue // try next NS } // Step 2: Build A query with RecursionDesired=false (authoritative query). fmt.Fprintf(w, "[dns] Stage 2: Querying %s (%s) for %s A (RD=0)\n", strings.TrimSuffix(nsHostname, "."), nsIP, hostname) id := uint16(rand.Uint32()) //nolint:gosec msg := dnsmessage.Message{ Header: dnsmessage.Header{ ID: id, RecursionDesired: false, }, Questions: []dnsmessage.Question{{ Name: qname, Type: dnsmessage.TypeA, Class: dnsmessage.ClassINET, }}, } packed, packErr := msg.Pack() if packErr != nil { return nil, "", "", fmt.Errorf("packing authoritative A query: %w", packErr) } // Step 3: Send via UDP transport. queryCtx, cancel := context.WithTimeout(ctx, timeout) resp, sendErr := UDPQuery(queryCtx, nsIP, packed, timeout) cancel() if sendErr != nil { continue // try next NS } // Step 4: Parse response. var parser dnsmessage.Parser respHeader, parseErr := parser.Start(resp) if parseErr != nil { continue } if respHeader.ID != id { continue } if respHeader.RCode == dnsmessage.RCodeRefused { continue // per R-000: treat Refused as "try next NS" } if respHeader.RCode != dnsmessage.RCodeSuccess { continue } if err := parser.SkipAllQuestions(); err != nil { continue } var foundIPs []string var foundCNAME string parseOK := true for { hdr, hdrErr := parser.AnswerHeader() if hdrErr == dnsmessage.ErrSectionDone { break } if hdrErr != nil { parseOK = false break } switch hdr.Type { case dnsmessage.TypeA: aRec, aErr := parser.AResource() if aErr != nil { parseOK = false break } foundIPs = append(foundIPs, net.IP(aRec.A[:]).String()) case dnsmessage.TypeCNAME: cnameRec, cErr := parser.CNAMEResource() if cErr != nil { parseOK = false break } foundCNAME = strings.TrimSuffix(cnameRec.CNAME.String(), ".") default: if skipErr := parser.SkipAnswer(); skipErr != nil { parseOK = false } } if !parseOK { break } } if !parseOK { continue } if len(foundIPs) > 0 { return foundIPs, "", nsHostname, nil } if foundCNAME != "" { return nil, foundCNAME, nsHostname, nil } // Empty response with RCodeSuccess — possible referral. Try next NS. } return nil, "", "", fmt.Errorf("all authoritative nameservers for zone %s are unreachable or returned empty responses", ns.Zone) } // resolveNSHostname resolves an NS hostname to an address string (IP or IP:port) // suitable for passing to UDPQuery. If the hostname is already an IP address or // IP:port string, it is returned directly (enabling test injection and handling // the rare case of IP addresses in NS records). Otherwise, each pool resolver is // queried for an A record and the first successful IP is returned. func resolveNSHostname(ctx context.Context, nsHostname string, resolvers []string, timeout time.Duration) (string, error) { // Strip trailing dot (NS records are FQDNs). nsHostname = strings.TrimSuffix(nsHostname, ".") // If the hostname is already an IP:port, use it directly. if h, _, err := net.SplitHostPort(nsHostname); err == nil { if net.ParseIP(h) != nil { return nsHostname, nil } } // If the hostname is a bare IP address, use it directly. if net.ParseIP(nsHostname) != nil { return nsHostname, nil } // Hostname — resolve via the pool. for _, r := range resolvers { queryCtx, cancel := context.WithTimeout(ctx, timeout) ips, err := queryA(queryCtx, r, nsHostname, timeout) cancel() if err == nil && len(ips) > 0 { return ips[0], nil } } return "", fmt.Errorf("could not resolve NS hostname %q via any resolver", nsHostname) }