Files

403 lines
12 KiB
Go

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
}
}
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)
}