// Package dns provides the low-level DNS query primitives used by ExploreDNS. // // It wraps the github.com/miekg/dns library to provide retrying, TCP fallback, // EDNS0 buffer size negotiation, and root server discovery. The package is // intentionally narrow: it sends iterative (non-recursive) queries and returns // the raw responses for the traversal engine to interpret. package dns import ( "context" "fmt" "net" "time" "github.com/miekg/dns" ) type QueryConfig struct { UDPSize int Timeout time.Duration Retries int UseTCP bool AllowTCP bool } func DefaultQueryConfig() *QueryConfig { return &QueryConfig{ UDPSize: DefaultEDNS0UDPSize(), Timeout: 5 * time.Second, Retries: 3, UseTCP: false, AllowTCP: true, } } type ExchangeFunc func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) func Query(ctx context.Context, server net.IP, name string, qtype uint16, cfg *QueryConfig) (*dns.Msg, error) { if cfg == nil { cfg = DefaultQueryConfig() } if cfg.UDPSize <= 0 { cfg.UDPSize = DefaultEDNS0UDPSize() } if cfg.Timeout <= 0 { cfg.Timeout = 5 * time.Second } return QueryWithExchange(ctx, server, name, qtype, cfg, realExchange) } func realExchange(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { addr := net.JoinHostPort(server, "53") var c *dns.Client if useTCP { c = &dns.Client{ Net: "tcp", ReadTimeout: 5 * time.Second, WriteTimeout: 5 * time.Second, } } else { c = &dns.Client{ Net: "udp", ReadTimeout: 5 * time.Second, WriteTimeout: 5 * 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", c.Net, addr, err) } return r, nil } func QueryWithExchange(ctx context.Context, server net.IP, name string, qtype uint16, cfg *QueryConfig, exchangeFn ExchangeFunc) (*dns.Msg, error) { if cfg == nil { cfg = DefaultQueryConfig() } if cfg.UDPSize <= 0 { cfg.UDPSize = DefaultEDNS0UDPSize() } msg := buildQuery(name, qtype, cfg.UDPSize) serverStr := server.String() var lastErr error for attempt := 0; attempt < cfg.Retries; attempt++ { if attempt > 0 { select { case <-ctx.Done(): return nil, fmt.Errorf("query retries cancelled: %w", ctx.Err()) case <-time.After(backoffDelay(attempt)): } } if cfg.UseTCP { resp, err := exchangeFn(ctx, serverStr, msg, true) if err != nil { lastErr = err continue } return resp, nil } resp, err := exchangeFn(ctx, serverStr, msg, false) if err != nil { lastErr = err continue } if resp.Truncated && cfg.AllowTCP { resp, err = exchangeFn(ctx, serverStr, msg, true) if err != nil { lastErr = err continue } return resp, nil } return resp, nil } return nil, fmt.Errorf("query %s %s failed after %d retries: %w", name, QNameType(qtype), cfg.Retries, lastErr) } func IterativeQuery(ctx context.Context, server net.IP, name string, qtype uint16, cfg *QueryConfig) (*dns.Msg, error) { if cfg == nil { cfg = DefaultQueryConfig() } if cfg.UDPSize <= 0 { cfg.UDPSize = DefaultEDNS0UDPSize() } return IterativeQueryWithExchange(ctx, server, name, qtype, cfg, realExchange) } func IterativeQueryWithExchange(ctx context.Context, server net.IP, name string, qtype uint16, cfg *QueryConfig, exchangeFn ExchangeFunc) (*dns.Msg, error) { if cfg == nil { cfg = DefaultQueryConfig() } if cfg.UDPSize <= 0 { cfg.UDPSize = DefaultEDNS0UDPSize() } msg := buildQuery(name, qtype, cfg.UDPSize) msg.RecursionDesired = false serverStr := server.String() var lastErr error for attempt := 0; attempt < cfg.Retries; attempt++ { if attempt > 0 { select { case <-ctx.Done(): return nil, fmt.Errorf("query retries cancelled: %w", ctx.Err()) case <-time.After(backoffDelay(attempt)): } } if cfg.UseTCP { resp, err := exchangeFn(ctx, serverStr, msg, true) if err != nil { lastErr = err continue } return resp, nil } resp, err := exchangeFn(ctx, serverStr, msg, false) if err != nil { lastErr = err continue } if resp == nil { lastErr = fmt.Errorf("nil response") continue } if resp.Truncated && cfg.AllowTCP { resp, err = exchangeFn(ctx, serverStr, msg, true) if err != nil { lastErr = err continue } return resp, nil } return resp, nil } return nil, fmt.Errorf("iterative query %s %s failed after %d retries: %w", name, QNameType(qtype), cfg.Retries, lastErr) } // backoffDelay computes the wait duration before the given retry attempt (1-indexed). // Delays: attempt=1 → 100ms, attempt=2 → 200ms, attempt=3 → 400ms, capped at 2s. func backoffDelay(attempt int) time.Duration { if attempt <= 0 { return 0 } delay := time.Duration(uint(1)< maxDelay { return maxDelay } return delay } func buildQuery(name string, qtype uint16, udpSize int) *dns.Msg { m := new(dns.Msg) m.SetQuestion(dns.Fqdn(name), qtype) m.RecursionDesired = true m.SetEdns0(uint16(udpSize), false) return m }