403 lines
12 KiB
Go
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)
|
|
}
|