package dns import ( "context" "fmt" "net" "time" "github.com/miekg/dns" ) type QueryConfig struct { UDPSize int Timeout time.Duration Retries int UseTCP bool } func DefaultQueryConfig() *QueryConfig { return &QueryConfig{ UDPSize: DefaultEDNS0UDPSize(), Timeout: 5 * time.Second, Retries: 3, UseTCP: false, } } 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(100 * time.Millisecond): } } 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 { 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 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 }