package resolver import ( "context" "fmt" "io" "math/rand" "net" "strings" "time" "golang.org/x/net/dns/dnsmessage" ) // ParallelAFallback sends A queries for hostname to all resolvers simultaneously. // Returns the first successful result (first response with at least one A record). // The buffered channel ensures goroutines never block on write even after the // collector returns early. // w receives Stage 3 diagnostic lines (pass io.Discard to suppress). func ParallelAFallback(ctx context.Context, w io.Writer, resolvers []string, hostname string, timeout time.Duration) ([]string, error) { if len(resolvers) == 0 { return nil, fmt.Errorf("no resolvers available for fallback A lookup of %s", hostname) } type aResult struct { Resolver string IPs []string Err error } ch := make(chan aResult, len(resolvers)) fmt.Fprintf(w, "[dns] Stage 3: Parallel A fallback for %s (%d resolvers)\n", hostname, len(resolvers)) for _, r := range resolvers { go func(resolver string) { queryCtx, cancel := context.WithTimeout(ctx, timeout) defer cancel() ips, err := queryA(queryCtx, resolver, hostname, timeout) ch <- aResult{Resolver: resolver, IPs: ips, Err: err} }(r) } // Collect all results to avoid goroutine leaks (channel is buffered). var lastErr error for i := 0; i < len(resolvers); i++ { result := <-ch if result.Err == nil && len(result.IPs) > 0 { fmt.Fprintf(w, "[dns] %s \u2192 %s\n", result.Resolver, strings.Join(result.IPs, ", ")) // First success wins; remaining goroutines write to the buffered channel // and exit cleanly even though we return early here. return result.IPs, nil } if result.Err != nil { errLabel := "error" if strings.Contains(result.Err.Error(), "NXDOMAIN") { errLabel = "NXDOMAIN" } fmt.Fprintf(w, "[dns] %s \u2192 %s\n", result.Resolver, errLabel) lastErr = result.Err } } if lastErr != nil { return nil, fmt.Errorf("all resolvers failed for %s: %w", hostname, lastErr) } return nil, fmt.Errorf("no resolver returned A records for %s", hostname) } // queryA sends a single recursive A query for hostname to resolver and returns // the IP addresses from the answer section. // NXDOMAIN is returned as an error. func queryA(ctx context.Context, resolver, hostname string, timeout time.Duration) ([]string, error) { fqdn := hostname if !strings.HasSuffix(fqdn, ".") { fqdn += "." } name, err := dnsmessage.NewName(fqdn) if err != nil { return nil, fmt.Errorf("invalid hostname %q: %w", hostname, err) } id := uint16(rand.Uint32()) //nolint:gosec msg := dnsmessage.Message{ Header: dnsmessage.Header{ ID: id, RecursionDesired: true, }, Questions: []dnsmessage.Question{{ Name: name, Type: dnsmessage.TypeA, Class: dnsmessage.ClassINET, }}, } packed, err := msg.Pack() if err != nil { return nil, fmt.Errorf("packing A query: %w", err) } resp, err := UDPQuery(ctx, resolver, packed, timeout) if err != nil { return nil, err } var parser dnsmessage.Parser respHeader, err := parser.Start(resp) if err != nil { return nil, fmt.Errorf("parsing A response: %w", err) } if respHeader.ID != id { return nil, fmt.Errorf("A response ID mismatch (expected %d, got %d)", id, respHeader.ID) } if respHeader.RCode == dnsmessage.RCodeNameError { return nil, fmt.Errorf("NXDOMAIN for %s", hostname) } if respHeader.RCode != dnsmessage.RCodeSuccess { return nil, fmt.Errorf("A query for %s returned %s", hostname, respHeader.RCode) } if err := parser.SkipAllQuestions(); err != nil { return nil, fmt.Errorf("skipping questions in A response: %w", err) } var ips []string for { hdr, err := parser.AnswerHeader() if err == dnsmessage.ErrSectionDone { break } if err != nil { return nil, fmt.Errorf("parsing A answer header: %w", err) } switch hdr.Type { case dnsmessage.TypeA: aRec, err := parser.AResource() if err != nil { return nil, fmt.Errorf("parsing A record: %w", err) } ips = append(ips, net.IP(aRec.A[:]).String()) default: if err := parser.SkipAnswer(); err != nil { return nil, fmt.Errorf("skipping A answer: %w", err) } } } if len(ips) == 0 { return nil, fmt.Errorf("no A records returned for %s", hostname) } return ips, nil }