package dns import ( "context" "fmt" "net" "time" "github.com/miekg/dns" ) // QueryConfig controls the transport-level behaviour of DNS queries. type QueryConfig struct { // UDPSize is the EDNS0 advertised UDP payload size in bytes. UDPSize int // Timeout is the per-attempt network timeout. Timeout time.Duration // Retries is the total number of query attempts (first try + retries - 1 extra attempts). Retries int // UseTCP forces all queries over TCP when true. UseTCP bool // AllowTCP enables automatic TCP retry when a UDP response is truncated. AllowTCP bool } // DefaultQueryConfig returns a QueryConfig with sensible defaults: // 2048-byte UDP buffer, 5-second timeout, 3 retries, UDP with TCP fallback. func DefaultQueryConfig() *QueryConfig { return &QueryConfig{ UDPSize: DefaultEDNS0UDPSize(), Timeout: 5 * time.Second, Retries: 3, UseTCP: false, AllowTCP: true, } } // ExchangeFunc is a function that sends a DNS message to server and returns // the response. Injecting an ExchangeFunc in tests avoids real network calls. type ExchangeFunc func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) // Query sends a standard (recursion-desired) DNS query to server for name/qtype // using the provided cfg. cfg may be nil; DefaultQueryConfig is used in that case. // Retries with exponential back-off are performed on transient errors. 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 } // QueryWithExchange is the testable core of Query. It uses exchangeFn instead // of the real network, enabling deterministic unit tests. 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) } // IterativeQuery sends a non-recursive (RD=false) DNS query intended for // authoritative nameservers. Use this during iterative traversal to prevent // resolvers from answering on behalf of the authoritative server. 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) } // IterativeQueryWithExchange is the testable core of IterativeQuery. 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 }