Refactor code structure for improved readability and maintainability
This commit is contained in:
@@ -0,0 +1,150 @@
|
||||
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
|
||||
}
|
||||
Reference in New Issue
Block a user