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") } } func TestIterativeQueryWithExchangeSuccess(t *testing.T) { answerResp := new(dns.Msg) answerResp.SetReply(new(dns.Msg)) answerResp.Answer = append(answerResp.Answer, &dns.A{ Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, A: net.ParseIP("1.2.3.4"), }) exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { if msg.RecursionDesired { t.Error("IterativeQuery should send RD=false") } return answerResp.Copy(), nil } server := net.ParseIP("198.41.0.4") resp, err := IterativeQueryWithExchange(context.Background(), server, "example.com", dns.TypeA, nil, exchangeFn) if err != nil { t.Fatalf("unexpected error: %v", err) } if len(resp.Answer) == 0 { t.Fatal("expected answer records") } } func TestIterativeQueryWithExchangeRetry(t *testing.T) { callCount := 0 answerResp := new(dns.Msg) answerResp.SetReply(new(dns.Msg)) answerResp.Answer = append(answerResp.Answer, &dns.A{ Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, A: net.ParseIP("1.2.3.4"), }) exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { callCount++ if callCount < 2 { return nil, errors.New("transient error") } return answerResp.Copy(), nil } cfg := &QueryConfig{UDPSize: 2048, Retries: 3, AllowTCP: true} server := net.ParseIP("198.41.0.4") resp, err := IterativeQueryWithExchange(context.Background(), server, "example.com", dns.TypeA, cfg, exchangeFn) if err != nil { t.Fatalf("unexpected error: %v", err) } if resp == nil { t.Fatal("expected non-nil response after retry") } } func TestIterativeQueryWithExchangeAllFail(t *testing.T) { exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { return nil, errors.New("server unreachable") } cfg := &QueryConfig{UDPSize: 2048, Retries: 2, AllowTCP: false} server := net.ParseIP("198.41.0.4") _, err := IterativeQueryWithExchange(context.Background(), server, "example.com", dns.TypeA, cfg, exchangeFn) if err == nil { t.Fatal("expected error when all attempts fail") } } func TestIterativeQueryWithExchangeTCPFallback(t *testing.T) { truncatedResp := new(dns.Msg) truncatedResp.SetReply(new(dns.Msg)) truncatedResp.Truncated = true 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("1.2.3.4"), }) exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { if !useTCP { return truncatedResp.Copy(), nil } return fullResp.Copy(), nil } cfg := &QueryConfig{UDPSize: 2048, Retries: 1, AllowTCP: true} server := net.ParseIP("198.41.0.4") resp, err := IterativeQueryWithExchange(context.Background(), server, "example.com", dns.TypeA, cfg, exchangeFn) if err != nil { t.Fatalf("unexpected error: %v", err) } if len(resp.Answer) == 0 { t.Fatal("expected answer after TCP fallback") } } func TestIterativeQueryWithExchangeContextCancelled(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) // Cancel the context immediately so the retry loop aborts during backoff cancel() exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { return nil, errors.New("error") } cfg := &QueryConfig{UDPSize: 2048, Retries: 5, AllowTCP: false} server := net.ParseIP("198.41.0.4") _, err := IterativeQueryWithExchange(ctx, server, "example.com", dns.TypeA, cfg, exchangeFn) if err == nil { t.Fatal("expected error when context cancelled") } } func TestIterativeQueryWithExchangeNilResponse(t *testing.T) { callCount := 0 exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { callCount++ return nil, nil // nil response, no error } cfg := &QueryConfig{UDPSize: 2048, Retries: 2, AllowTCP: false} server := net.ParseIP("198.41.0.4") _, err := IterativeQueryWithExchange(context.Background(), server, "example.com", dns.TypeA, cfg, exchangeFn) if err == nil { t.Fatal("expected error for nil responses") } } func TestIterativeQueryWithExchangeUseTCP(t *testing.T) { var wasTCP bool answerResp := new(dns.Msg) answerResp.SetReply(new(dns.Msg)) answerResp.Answer = append(answerResp.Answer, &dns.A{ Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, A: net.ParseIP("1.2.3.4"), }) exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { wasTCP = useTCP return answerResp.Copy(), nil } cfg := &QueryConfig{UDPSize: 2048, Retries: 1, UseTCP: true} server := net.ParseIP("198.41.0.4") _, err := IterativeQueryWithExchange(context.Background(), server, "example.com", dns.TypeA, cfg, exchangeFn) if err != nil { t.Fatalf("unexpected error: %v", err) } if !wasTCP { t.Error("expected TCP exchange when UseTCP=true") } }