Files

151 lines
4.2 KiB
Go

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
}