package dns import ( "context" "errors" "net" "sync" "testing" "github.com/miekg/dns" ) func TestIterativeQueryWithExchangeSuccess(t *testing.T) { resp := new(dns.Msg) resp.SetReply(new(dns.Msg)) resp.RecursionDesired = false resp.Answer = append(resp.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) { if msg.RecursionDesired { t.Error("iterative query should have RecursionDesired=false") } return resp.Copy(), nil } cfg := &QueryConfig{ UDPSize: 2048, Retries: 1, UseTCP: false, } server := net.ParseIP("8.8.8.8") result, err := IterativeQueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn) if err != nil { t.Fatalf("unexpected error: %v", err) } if len(result.Answer) != 1 { t.Fatalf("expected 1 answer, got %d", len(result.Answer)) } } func TestIterativeQueryWithExchangeNilConfig(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 := IterativeQueryWithExchange(context.Background(), server, "example.com", TypeA, nil, exchangeFn) if err != nil { t.Fatalf("unexpected error with nil config: %v", err) } } func TestIterativeQueryWithExchangeNilResponse(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, nil } cfg := &QueryConfig{ UDPSize: 2048, Retries: 2, UseTCP: false, } server := net.ParseIP("8.8.8.8") _, err := IterativeQueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn) if err == nil { t.Fatal("expected error when response is nil") } if callCount != 2 { t.Errorf("expected 2 attempts, got %d", callCount) } } func TestIterativeQueryWithExchangeTCPFallback(t *testing.T) { truncated := new(dns.Msg) truncated.Truncated = true truncated.SetReply(new(dns.Msg)) full := new(dns.Msg) full.SetReply(new(dns.Msg)) full.Answer = append(full.Answer, &dns.A{ Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA}, A: net.ParseIP("1.2.3.4"), }) var mu sync.Mutex calls := []bool{} exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { mu.Lock() calls = append(calls, useTCP) mu.Unlock() if !useTCP { return truncated.Copy(), nil } return full.Copy(), nil } cfg := &QueryConfig{ UDPSize: 2048, Retries: 1, UseTCP: false, AllowTCP: true, } server := net.ParseIP("8.8.8.8") result, err := IterativeQueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn) if err != nil { t.Fatalf("unexpected error: %v", err) } if len(calls) != 2 { t.Errorf("expected 2 calls (UDP + TCP), got %d", len(calls)) } if len(result.Answer) != 1 { t.Fatalf("expected 1 answer, got %d", len(result.Answer)) } } func TestIterativeQueryWithExchangeAlwaysTCP(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() calls = append(calls, useTCP) mu.Unlock() return resp.Copy(), nil } cfg := &QueryConfig{ UDPSize: 2048, Retries: 1, UseTCP: true, } server := net.ParseIP("8.8.8.8") _, err := IterativeQueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn) if err != nil { t.Fatalf("unexpected error: %v", err) } if len(calls) != 1 || !calls[0] { t.Errorf("expected single TCP call, got %v", calls) } } func TestIterativeQueryWithExchangeRetriesOnFailure(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, Retries: 3, UseTCP: false, } server := net.ParseIP("8.8.8.8") _, err := IterativeQueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn) if err == nil { t.Fatal("expected error") } if callCount != 3 { t.Errorf("expected 3 calls, got %d", callCount) } } func TestIterativeQueryWithExchangeContextCancel(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, Retries: 1, UseTCP: false, } server := net.ParseIP("8.8.8.8") _, err := IterativeQueryWithExchange(ctx, server, "example.com", TypeA, cfg, exchangeFn) if err == nil { t.Fatal("expected error on cancelled context") } } func TestExtractNSNames(t *testing.T) { rrs := []dns.RR{ &dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS}, Ns: "a.root-servers.net."}, &dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS}, Ns: "b.root-servers.net."}, } names := extractNSNames(rrs) if len(names) != 2 { t.Fatalf("expected 2 names, got %d", len(names)) } } func TestExtractNSNamesEmpty(t *testing.T) { names := extractNSNames(nil) if len(names) != 0 { t.Errorf("expected 0 names, got %d", len(names)) } } func TestIterativeQueryWithExchangeZeroUDPSize(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, // should use default Retries: 1, } server := net.ParseIP("8.8.8.8") _, err := IterativeQueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn) if err != nil { t.Fatalf("unexpected error: %v", err) } }