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

332 lines
8.1 KiB
Go

package dns
import (
"context"
"errors"
"net"
"sync"
"testing"
"github.com/miekg/dns"
)
func TestDefaultQueryConfig(t *testing.T) {
cfg := DefaultQueryConfig()
if cfg == nil {
t.Fatal("DefaultQueryConfig returned nil")
}
if cfg.UDPSize != 2048 {
t.Errorf("UDPSize = %d, want 2048", cfg.UDPSize)
}
if cfg.Retries != 3 {
t.Errorf("Retries = %d, want 3", cfg.Retries)
}
if cfg.UseTCP {
t.Error("UseTCP should be false by default")
}
}
func TestBuildQuery(t *testing.T) {
msg := buildQuery("example.com.", TypeA, 2048)
if len(msg.Question) != 1 {
t.Fatalf("expected 1 question, got %d", len(msg.Question))
}
q := msg.Question[0]
if q.Name != "example.com." {
t.Errorf("question name = %q, want %q", q.Name, "example.com.")
}
if q.Qtype != TypeA {
t.Errorf("question type = %d, want %d", q.Qtype, TypeA)
}
if !msg.RecursionDesired {
t.Error("RecursionDesired should be true")
}
if opt := msg.IsEdns0(); opt == nil {
t.Error("expected EDNS0 OPT record")
} else if opt.UDPSize() != 2048 {
t.Errorf("EDNS0 UDPSize = %d, want 2048", opt.UDPSize())
}
}
func TestBuildQueryFqdn(t *testing.T) {
msg := buildQuery("example.com", TypeA, 4096)
q := msg.Question[0]
if q.Name != "example.com." {
t.Errorf("Fqdn not applied: got %q, want %q", q.Name, "example.com.")
}
}
func TestQueryWithExchangeSuccess(t *testing.T) {
expectedResp := new(dns.Msg)
expectedResp.SetReply(new(dns.Msg))
expectedResp.Answer = append(expectedResp.Answer, &dns.A{
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300},
A: net.ParseIP("93.184.216.34"),
})
exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
return expectedResp.Copy(), nil
}
cfg := &QueryConfig{
UDPSize: 2048,
Timeout: 5,
Retries: 1,
UseTCP: false,
}
server := net.ParseIP("8.8.8.8")
resp, err := QueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if len(resp.Answer) != 1 {
t.Fatalf("expected 1 answer, got %d", len(resp.Answer))
}
}
func TestQueryTCPFallbackOnTruncation(t *testing.T) {
truncatedResp := new(dns.Msg)
truncatedResp.Truncated = true
truncatedResp.SetReply(new(dns.Msg))
fullResp := new(dns.Msg)
fullResp.SetReply(new(dns.Msg))
fullResp.Answer = append(fullResp.Answer, &dns.A{
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300},
A: net.ParseIP("93.184.216.34"),
})
var mu sync.Mutex
calls := []bool{}
exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
mu.Lock()
defer mu.Unlock()
calls = append(calls, useTCP)
if !useTCP {
return truncatedResp.Copy(), nil
}
return fullResp.Copy(), nil
}
cfg := &QueryConfig{
UDPSize: 2048,
Timeout: 5,
Retries: 1,
UseTCP: false,
AllowTCP: true,
}
server := net.ParseIP("8.8.8.8")
resp, err := QueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if len(calls) != 2 {
t.Fatalf("expected 2 exchange calls (UDP then TCP), got %d", len(calls))
}
if calls[0] != false {
t.Error("first call should be UDP")
}
if calls[1] != true {
t.Error("second call should be TCP")
}
if len(resp.Answer) != 1 {
t.Fatalf("expected 1 answer from TCP fallback, got %d", len(resp.Answer))
}
}
func TestQueryAlwaysTCP(t *testing.T) {
resp := new(dns.Msg)
resp.SetReply(new(dns.Msg))
var mu sync.Mutex
calls := []bool{}
exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
mu.Lock()
defer mu.Unlock()
calls = append(calls, useTCP)
return resp.Copy(), nil
}
cfg := &QueryConfig{
UDPSize: 2048,
Timeout: 5,
Retries: 1,
UseTCP: true,
}
server := net.ParseIP("8.8.8.8")
_, err := QueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if len(calls) != 1 {
t.Fatalf("expected 1 exchange call, got %d", len(calls))
}
if !calls[0] {
t.Error("expected TCP call when UseTCP is true")
}
}
func TestQueryRetriesOnFailure(t *testing.T) {
var mu sync.Mutex
callCount := 0
exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
mu.Lock()
callCount++
mu.Unlock()
return nil, errors.New("connection refused")
}
cfg := &QueryConfig{
UDPSize: 2048,
Timeout: 5,
Retries: 3,
UseTCP: false,
}
server := net.ParseIP("8.8.8.8")
_, err := QueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn)
if err == nil {
t.Fatal("expected error after retries exhausted")
}
if callCount != 3 {
t.Errorf("expected 3 calls (retries exhausted), got %d", callCount)
}
}
func TestQueryContextCancellation(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
cancel()
exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
return nil, ctx.Err()
}
cfg := &QueryConfig{
UDPSize: 2048,
Timeout: 5,
Retries: 1,
UseTCP: false,
}
server := net.ParseIP("8.8.8.8")
_, err := QueryWithExchange(ctx, server, "example.com", TypeA, cfg, exchangeFn)
if err == nil {
t.Fatal("expected error on cancelled context")
}
}
func TestQueryNilConfigUsesDefaults(t *testing.T) {
resp := new(dns.Msg)
resp.SetReply(new(dns.Msg))
exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
return resp.Copy(), nil
}
server := net.ParseIP("8.8.8.8")
_, err := QueryWithExchange(context.Background(), server, "example.com", TypeA, nil, exchangeFn)
if err != nil {
t.Fatalf("unexpected error with nil config: %v", err)
}
}
func TestQueryZeroValuesUseDefaults(t *testing.T) {
resp := new(dns.Msg)
resp.SetReply(new(dns.Msg))
exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
return resp.Copy(), nil
}
cfg := &QueryConfig{
UDPSize: 0,
Timeout: 0,
Retries: 1,
UseTCP: false,
}
server := net.ParseIP("8.8.8.8")
_, err := QueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
}
func TestQueryTCPFallbackFailsThenRetries(t *testing.T) {
truncatedResp := new(dns.Msg)
truncatedResp.Truncated = true
truncatedResp.SetReply(new(dns.Msg))
var mu sync.Mutex
callCount := 0
exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
mu.Lock()
callCount++
mu.Unlock()
if !useTCP {
return truncatedResp.Copy(), nil
}
return nil, errors.New("tcp failed")
}
cfg := &QueryConfig{
UDPSize: 2048,
Timeout: 5,
Retries: 2,
UseTCP: false,
AllowTCP: true,
}
server := net.ParseIP("8.8.8.8")
_, err := QueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn)
if err == nil {
t.Fatal("expected error when TCP fallback always fails")
}
if callCount != 4 {
t.Errorf("expected 4 calls (2 retries x UDP+TCP), got %d", callCount)
}
}
func TestQueryNoTCPFallbackWhenDisabled(t *testing.T) {
truncatedResp := new(dns.Msg)
truncatedResp.Truncated = true
truncatedResp.SetReply(new(dns.Msg))
var mu sync.Mutex
calls := []bool{}
exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
mu.Lock()
defer mu.Unlock()
calls = append(calls, useTCP)
return truncatedResp.Copy(), nil
}
cfg := &QueryConfig{
UDPSize: 2048,
Timeout: 5,
Retries: 1,
UseTCP: false,
AllowTCP: false, // TCP fallback must be suppressed
}
server := net.ParseIP("8.8.8.8")
resp, err := QueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
// Only one UDP call; no TCP fallback.
if len(calls) != 1 {
t.Fatalf("expected 1 exchange call (no TCP fallback), got %d", len(calls))
}
if calls[0] != false {
t.Error("expected UDP-only call")
}
if !resp.Truncated {
t.Error("expected truncated response to be returned as-is")
}
}