diff --git a/internal/dns/iterative_test.go b/internal/dns/iterative_test.go index 2f506ce..854064c 100644 --- a/internal/dns/iterative_test.go +++ b/internal/dns/iterative_test.go @@ -10,38 +10,6 @@ import ( "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)) @@ -57,76 +25,6 @@ func TestIterativeQueryWithExchangeNilConfig(t *testing.T) { } } -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)) @@ -203,24 +101,6 @@ func TestIterativeQueryWithExchangeContextCancel(t *testing.T) { } } -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)) diff --git a/internal/output/coverage_test.go b/internal/output/coverage_test.go index f2587b7..d756216 100644 --- a/internal/output/coverage_test.go +++ b/internal/output/coverage_test.go @@ -52,18 +52,6 @@ func TestRRDataString(t *testing.T) { } } -func TestRRDataStringDefault(t *testing.T) { - soa := &miekgdns.SOA{ - Hdr: miekgdns.RR_Header{Name: "example.com.", Rrtype: miekgdns.TypeSOA, Class: miekgdns.ClassINET, Ttl: 3600}, - Ns: "ns1.example.com.", - Mbox: "admin.example.com.", - } - got := rrDataString(soa) - if got == "" { - t.Error("expected non-empty string for SOA default case") - } -} - func TestSummaryTypeLabel(t *testing.T) { cases := []struct { input string @@ -87,22 +75,6 @@ func TestSummaryTypeLabel(t *testing.T) { } } -func TestContainsString(t *testing.T) { - items := []string{"a", "b", "c"} - if !containsString(items, "a") { - t.Error("containsString should find 'a'") - } - if !containsString(items, "c") { - t.Error("containsString should find 'c'") - } - if containsString(items, "d") { - t.Error("containsString should not find 'd'") - } - if containsString(nil, "a") { - t.Error("containsString on nil should return false") - } -} - func TestCollectServers(t *testing.T) { ref := traverse.NewReferral("example.com.", dns.TypeA, "com.", 1, 0.5, nil) resp := &traverse.Response{ @@ -848,20 +820,6 @@ func TestJSONFormatterWriteProgressShowProgressFalse(t *testing.T) { // ---- runner.go coverage ---- -func TestCollectUniqueServerIPs(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - results := []traverse.TraversalResult{ - {Referral: ref, Response: &traverse.Response{Server: net.ParseIP("1.2.3.4"), Referral: ref, Type: traverse.RespAnswer}}, - {Referral: ref, Response: &traverse.Response{Server: net.ParseIP("1.2.3.4"), Referral: ref, Type: traverse.RespAnswer}}, // duplicate - {Referral: ref, Response: &traverse.Response{Server: net.ParseIP("5.6.7.8"), Referral: ref, Type: traverse.RespAnswer}}, - {Referral: ref, Response: nil}, - } - ips := collectUniqueServerIPs(results) - if len(ips) != 2 { - t.Errorf("expected 2 unique IPs, got %d", len(ips)) - } -} - func TestRunTraversalNilTraverser(t *testing.T) { _, err := RunTraversal(context.Background(), nil, nil, nil, "example.com") if err == nil { @@ -876,13 +834,6 @@ func TestNewFormatterNilConfig(t *testing.T) { } } -func TestNewFormatterNilWriter(t *testing.T) { - f := NewFormatter(DefaultConfig(), nil) - if f == nil { - t.Fatal("NewFormatter with nil writer should not return nil") - } -} - func TestAttachHooksNilCfg(t *testing.T) { h := AttachHooks(nil, nil) if h != nil {