// Package dns provides the low-level DNS query primitives used by ExploreDNS. // // It wraps the github.com/miekg/dns library behind a single query path // (Client) that sends non-recursive (RD=0) queries with retrying, TCP // fallback on truncation, EDNS0 negotiation and a per-run packet cache, plus // root server discovery. Production uses the real wire exchange; tests inject // a mock ExchangeFunc into the exact same path. package dns import ( "context" "errors" "fmt" "net" "sync" "time" "github.com/miekg/dns" ) // QueryConfig controls transport parameters for the single query path. type QueryConfig struct { // UDPSize is the EDNS0 UDP payload size. An OPT record is only attached // when UDPSize > 512, mirroring dnstraverse's caching_resolver.rb. UDPSize int // Timeout is the per-attempt packet timeout (dnsruby packet_timeout, // dnstraverse default 2s). Timeout time.Duration // Retries is the total number of send attempts, matching dnsruby // retry_times: Resolver#generate_timeouts schedules retry_times // transmissions in total — the first immediately and retry k at // retry_delay*2^k seconds after the first. Values below 1 are clamped to // 1 so exactly one query is still sent (dnsruby with retry_times 0 would // send nothing and hang; this also fixes the old // "failed after 0 retries: %!w()" error). Retries int // RetryDelay is dnsruby's retry_delay (dnstraverse default 2s). RetryDelay time.Duration // UseTCP forces every query over TCP (--always-tcp). UseTCP bool // AllowTCP enables the UDP→TCP retry when a response is truncated. AllowTCP bool } func DefaultQueryConfig() *QueryConfig { return &QueryConfig{ UDPSize: DefaultEDNS0UDPSize(), Timeout: 2 * time.Second, Retries: 2, RetryDelay: 2 * time.Second, UseTCP: false, AllowTCP: true, } } // withDefaults returns a copy of cfg with zero values replaced by defaults. func (cfg *QueryConfig) withDefaults() *QueryConfig { if cfg == nil { return DefaultQueryConfig() } out := *cfg if out.UDPSize <= 0 { out.UDPSize = DefaultEDNS0UDPSize() } if out.Timeout <= 0 { out.Timeout = 2 * time.Second } if out.Retries < 1 { out.Retries = 1 } if out.RetryDelay <= 0 { out.RetryDelay = 2 * time.Second } return &out } // ExchangeFunc performs one wire exchange. server is either a bare host/IP // (port 53 implied) or an explicit host:port. type ExchangeFunc func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) func realExchange(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { addr := server if _, _, err := net.SplitHostPort(server); err != nil { addr = net.JoinHostPort(server, "53") } proto := "udp" if useTCP { proto = "tcp" } c := &dns.Client{ Net: proto, ReadTimeout: 2 * time.Second, WriteTimeout: 2 * time.Second, } if deadline, ok := ctx.Deadline(); ok { c.ReadTimeout = time.Until(deadline) c.WriteTimeout = time.Until(deadline) } r, _, err := c.ExchangeContext(ctx, msg, addr) if err != nil { return nil, fmt.Errorf("dns exchange (%s) with %s: %w", proto, addr, err) } return r, nil } // Client is the single query path used identically by production and tests. // Every query is non-recursive (RD=0) and deduplicated by a per-run packet // cache keyed (server IP, qname, qclass, qtype, udpsize), mirroring // dnstraverse's caching_resolver.rb: repeat askers replay the cached answer // (or cached failure) without touching the wire. type Client struct { cfg *QueryConfig exchange ExchangeFunc mu sync.Mutex cache map[packetKey]*packetEntry requests int cacheHits int } type packetKey struct { server string qname string qclass uint16 qtype uint16 udpsize int } type packetEntry struct { once sync.Once msg *dns.Msg err error } // NewClient creates a Client. A nil exchange means the real wire exchange; // tests pass a mock so no packets leave the process. func NewClient(cfg *QueryConfig, exchange ExchangeFunc) *Client { if exchange == nil { exchange = realExchange } return &Client{ cfg: cfg.withDefaults(), exchange: exchange, cache: make(map[packetKey]*packetEntry), } } // Requests reports how many queries were asked of the client (cache hits included). func (c *Client) Requests() int { c.mu.Lock() defer c.mu.Unlock() return c.requests } // CacheHits reports how many queries were served from the packet cache. func (c *Client) CacheHits() int { c.mu.Lock() defer c.mu.Unlock() return c.cacheHits } // Query sends a non-recursive query for name/qtype (class IN) to server and // returns the response plus any warnings gathered along the way (EDNS0 // fallback, recursion offered, truncation). A non-nil error corresponds to // dnstraverse's "exception" status (network failure after all retries). func (c *Client) Query(ctx context.Context, server net.IP, name string, qtype uint16) (*dns.Msg, []string, error) { msg, err := c.cachedExchange(ctx, server, name, qtype, c.cfg.UDPSize) if err != nil { return nil, nil, err } var warnings []string // EDNS0 fallback (decoded_query.rb makequery_message): FORMERR/NOTIMP/ // SERVFAIL with udpsize > 512 may mean the server chokes on OPT; retry // once at 512 and keep the retry only if it clears the error. if c.cfg.UDPSize > MinEDNS0UDPSize() && ednsFailure(msg.Rcode) { retryMsg, retryErr := c.cachedExchange(ctx, server, name, qtype, MinEDNS0UDPSize()) if retryErr == nil && !ednsFailure(retryMsg.Rcode) { warnings = append(warnings, fmt.Sprintf("%s doesn't seem to support EDNS0", server)) msg = retryMsg } } // msg_comment with want_recursion=false (message_utility.rb). if msg.RecursionAvailable { warnings = append(warnings, fmt.Sprintf("%s allows recursion", server)) } if msg.Truncated { warnings = append(warnings, fmt.Sprintf("%s sent truncated packet", server)) } return msg, warnings, nil } func ednsFailure(rcode int) bool { return rcode == dns.RcodeFormatError || rcode == dns.RcodeNotImplemented || rcode == dns.RcodeServerFailure } // cachedExchange sends at most one wire query per packet cache key; the // outcome (response or error) is cached and replayed for repeat askers. func (c *Client) cachedExchange(ctx context.Context, server net.IP, name string, qtype uint16, udpsize int) (*dns.Msg, error) { key := packetKey{ server: server.String(), qname: dns.CanonicalName(name), qclass: dns.ClassINET, qtype: qtype, udpsize: udpsize, } c.mu.Lock() c.requests++ entry, ok := c.cache[key] if ok { c.cacheHits++ } else { entry = &packetEntry{} c.cache[key] = entry } c.mu.Unlock() entry.once.Do(func() { entry.msg, entry.err = exchangeWithRetry(ctx, c.exchange, key.server, buildQuery(name, qtype, udpsize), c.cfg) }) return copyMsg(entry.msg), entry.err } // exchangeWithRetry implements dnsruby's retry schedule (resolver.rb // generate_timeouts): cfg.Retries is the TOTAL number of transmissions — the // first goes immediately and retry k is sent retry_delay*2^k seconds after // the first (so gaps of 2d, 2d, 4d, 8d, ...). Each attempt gets its own // cfg.Timeout (dnsruby packet_timeout). A truncated UDP reply is retried over // TCP within the same attempt when cfg.AllowTCP. func exchangeWithRetry(ctx context.Context, exchange ExchangeFunc, server string, msg *dns.Msg, cfg *QueryConfig) (*dns.Msg, error) { attempts := cfg.Retries if attempts < 1 { attempts = 1 } var lastErr error for attempt := 0; attempt < attempts; attempt++ { if attempt > 0 { select { case <-ctx.Done(): return nil, fmt.Errorf("query retries cancelled: %w", ctx.Err()) case <-time.After(retryGap(cfg.RetryDelay, attempt)): } } resp, err := exchangeOnce(ctx, exchange, server, msg, cfg) if err != nil { lastErr = err continue } return resp, nil } q := msg.Question[0] return nil, fmt.Errorf("query %s %s to %s failed after %d attempts: %w", q.Name, QNameType(q.Qtype), server, attempts, lastErr) } // retryGap returns the wait before retry number `retry` (1-based). dnsruby // sends retry k at absolute time retry_delay*2^k, so the gap is 2d before the // first retry and d*2^(k-1) for each retry after that. func retryGap(d time.Duration, retry int) time.Duration { if retry <= 1 { return 2 * d } return d << uint(retry-1) } func exchangeOnce(ctx context.Context, exchange ExchangeFunc, server string, msg *dns.Msg, cfg *QueryConfig) (*dns.Msg, error) { actx, cancel := context.WithTimeout(ctx, cfg.Timeout) defer cancel() resp, err := exchange(actx, server, msg, cfg.UseTCP) if err != nil { return nil, err } if resp == nil { return nil, errors.New("nil response") } if !cfg.UseTCP && resp.Truncated && cfg.AllowTCP { resp, err = exchange(actx, server, msg, true) if err != nil { return nil, err } if resp == nil { return nil, errors.New("nil response") } } return resp, nil } // buildQuery constructs a non-recursive (RD=0) class IN query. The EDNS0 OPT // record is attached only when udpsize > 512 (caching_resolver.rb adds OPT // under the same condition), with the DO bit off. func buildQuery(name string, qtype uint16, udpsize int) *dns.Msg { m := new(dns.Msg) m.SetQuestion(dns.Fqdn(name), qtype) m.RecursionDesired = false if udpsize > MinEDNS0UDPSize() { m.SetEdns0(uint16(udpsize), false) } return m } func copyMsg(m *dns.Msg) *dns.Msg { if m == nil { return nil } return m.Copy() }