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