151 lines
4.2 KiB
Go
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
|
|
}
|