Files
ExploreDNS/internal/dns/query.go
T
2026-06-07 16:08:39 +00:00

205 lines
4.4 KiB
Go

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(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 && 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(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 == 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)
}
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
}