package dns import ( "context" "errors" "net" "strings" "sync" "testing" "time" "github.com/miekg/dns" ) // testClient returns a Client with fast retries wired to fn. func testClient(cfg *QueryConfig, fn ExchangeFunc) *Client { if cfg == nil { cfg = DefaultQueryConfig() } cfg.Timeout = time.Second cfg.RetryDelay = time.Millisecond return NewClient(cfg, fn) } func answerMsg(name string, ip string) *dns.Msg { m := new(dns.Msg) m.SetReply(new(dns.Msg)) m.Answer = append(m.Answer, &dns.A{ Hdr: dns.RR_Header{Name: name, Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, A: net.ParseIP(ip), }) return m } func TestDefaultQueryConfig(t *testing.T) { cfg := DefaultQueryConfig() if cfg.UDPSize != 2048 { t.Errorf("UDPSize = %d, want 2048", cfg.UDPSize) } if cfg.Retries != 2 { t.Errorf("Retries = %d, want 2", cfg.Retries) } if cfg.Timeout != 2*time.Second { t.Errorf("Timeout = %v, want 2s", cfg.Timeout) } if cfg.RetryDelay != 2*time.Second { t.Errorf("RetryDelay = %v, want 2s", cfg.RetryDelay) } if cfg.UseTCP { t.Error("UseTCP should be false by default") } if !cfg.AllowTCP { t.Error("AllowTCP should be true by default") } } func TestBuildQueryRDZeroWithEDNS(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("Fqdn not applied: got %q", q.Name) } if q.Qtype != TypeA { t.Errorf("question type = %d, want %d", q.Qtype, TypeA) } if msg.RecursionDesired { t.Error("RecursionDesired must be false on the query path") } opt := msg.IsEdns0() if opt == nil { t.Fatal("expected EDNS0 OPT record for udpsize > 512") } if opt.UDPSize() != 2048 { t.Errorf("EDNS0 UDPSize = %d, want 2048", opt.UDPSize()) } } func TestBuildQueryNoEDNSAt512(t *testing.T) { msg := buildQuery("example.com.", TypeA, 512) if msg.IsEdns0() != nil { t.Error("no OPT record should be attached when udpsize <= 512") } } func TestClientQuerySuccess(t *testing.T) { c := testClient(nil, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { if msg.RecursionDesired { t.Error("query must be sent with RD=0") } return answerMsg("example.com.", "93.184.216.34"), nil }) resp, warnings, err := c.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA) if err != nil { t.Fatalf("unexpected error: %v", err) } if len(resp.Answer) != 1 { t.Fatalf("expected 1 answer, got %d", len(resp.Answer)) } if len(warnings) != 0 { t.Errorf("unexpected warnings: %v", warnings) } } func TestClientPacketCache(t *testing.T) { var calls int c := testClient(nil, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { calls++ return answerMsg("example.com.", "1.2.3.4"), nil }) ctx := context.Background() server := net.ParseIP("8.8.8.8") for i := 0; i < 3; i++ { if _, _, err := c.Query(ctx, server, "example.com", TypeA); err != nil { t.Fatalf("query %d: %v", i, err) } } if calls != 1 { t.Errorf("expected 1 wire query for repeat askers, got %d", calls) } if c.Requests() != 3 { t.Errorf("Requests = %d, want 3", c.Requests()) } if c.CacheHits() != 2 { t.Errorf("CacheHits = %d, want 2", c.CacheHits()) } // Case differences must hit the same cache entry. if _, _, err := c.Query(ctx, server, "EXAMPLE.COM.", TypeA); err != nil { t.Fatal(err) } if calls != 1 { t.Errorf("case-insensitive lookup should hit cache, got %d wire calls", calls) } } func TestClientPacketCacheDistinctKeys(t *testing.T) { var calls int c := testClient(nil, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { calls++ return answerMsg("example.com.", "1.2.3.4"), nil }) ctx := context.Background() _, _, _ = c.Query(ctx, net.ParseIP("8.8.8.8"), "example.com", TypeA) _, _, _ = c.Query(ctx, net.ParseIP("9.9.9.9"), "example.com", TypeA) _, _, _ = c.Query(ctx, net.ParseIP("8.8.8.8"), "example.com", TypeAAAA) _, _, _ = c.Query(ctx, net.ParseIP("8.8.8.8"), "other.com", TypeA) if calls != 4 { t.Errorf("expected 4 wire queries for 4 distinct keys, got %d", calls) } } func TestClientPacketCacheReplaysErrors(t *testing.T) { var calls int c := testClient(&QueryConfig{Retries: 1}, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { calls++ return nil, errors.New("connection refused") }) ctx := context.Background() server := net.ParseIP("8.8.8.8") _, _, err1 := c.Query(ctx, server, "example.com", TypeA) _, _, err2 := c.Query(ctx, server, "example.com", TypeA) if err1 == nil || err2 == nil { t.Fatal("expected errors") } if calls != 1 { t.Errorf("failures must be cached too: got %d wire calls", calls) } } func TestClientEDNSFallback(t *testing.T) { var sizes []int c := testClient(nil, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { size := 512 if opt := msg.IsEdns0(); opt != nil { size = int(opt.UDPSize()) } sizes = append(sizes, size) m := new(dns.Msg) m.SetReply(msg) if size > 512 { m.Rcode = dns.RcodeFormatError return m, nil } m.Answer = append(m.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"), }) return m, nil }) resp, warnings, err := c.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA) if err != nil { t.Fatalf("unexpected error: %v", err) } if len(sizes) != 2 || sizes[0] != 2048 || sizes[1] != 512 { t.Fatalf("expected 2048 then 512 queries, got %v", sizes) } if resp.Rcode != dns.RcodeSuccess || len(resp.Answer) != 1 { t.Error("expected the 512-byte retry response to be returned") } want := "8.8.8.8 doesn't seem to support EDNS0" if len(warnings) != 1 || warnings[0] != want { t.Errorf("warnings = %v, want [%q]", warnings, want) } } func TestClientEDNSFallbackKeepsOriginalWhenRetryFails(t *testing.T) { c := testClient(nil, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { m := new(dns.Msg) m.SetReply(msg) m.Rcode = dns.RcodeServerFailure return m, nil }) resp, warnings, err := c.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA) if err != nil { t.Fatalf("unexpected error: %v", err) } if resp.Rcode != dns.RcodeServerFailure { t.Errorf("expected original SERVFAIL response, got rcode %d", resp.Rcode) } if len(warnings) != 0 { t.Errorf("no EDNS0 warning expected when the retry also fails: %v", warnings) } } func TestClientNoEDNSFallbackAt512(t *testing.T) { var calls int c := testClient(&QueryConfig{UDPSize: 512, Retries: 1}, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { calls++ m := new(dns.Msg) m.SetReply(msg) m.Rcode = dns.RcodeFormatError return m, nil }) resp, _, err := c.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA) if err != nil { t.Fatalf("unexpected error: %v", err) } if calls != 1 { t.Errorf("expected no fallback query at udpsize 512, got %d calls", calls) } if resp.Rcode != dns.RcodeFormatError { t.Errorf("expected FORMERR passthrough, got %d", resp.Rcode) } } func TestClientRecursionAvailableWarning(t *testing.T) { c := testClient(nil, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { m := answerMsg("example.com.", "1.2.3.4") m.RecursionAvailable = true return m, nil }) _, warnings, err := c.Query(context.Background(), net.ParseIP("192.0.2.1"), "example.com", TypeA) if err != nil { t.Fatal(err) } want := "192.0.2.1 allows recursion" if len(warnings) != 1 || warnings[0] != want { t.Errorf("warnings = %v, want [%q]", warnings, want) } } func TestClientTruncationWarningWhenTCPDisallowed(t *testing.T) { c := testClient(&QueryConfig{AllowTCP: false, Retries: 1}, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { if useTCP { t.Error("TCP must not be used when AllowTCP is false") } m := answerMsg("example.com.", "1.2.3.4") m.Truncated = true return m, nil }) resp, warnings, err := c.Query(context.Background(), net.ParseIP("192.0.2.1"), "example.com", TypeA) if err != nil { t.Fatal(err) } if !resp.Truncated { t.Error("expected truncated response to be returned as-is") } want := "192.0.2.1 sent truncated packet" if len(warnings) != 1 || warnings[0] != want { t.Errorf("warnings = %v, want [%q]", warnings, want) } } func TestClientTCPFallbackOnTruncation(t *testing.T) { var mu sync.Mutex var calls []bool c := testClient(nil, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { mu.Lock() calls = append(calls, useTCP) mu.Unlock() if !useTCP { m := new(dns.Msg) m.SetReply(msg) m.Truncated = true return m, nil } return answerMsg("example.com.", "93.184.216.34"), nil }) resp, warnings, err := c.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA) if err != nil { t.Fatalf("unexpected error: %v", err) } if len(calls) != 2 || calls[0] != false || calls[1] != true { t.Fatalf("expected UDP then TCP, got %v", calls) } if len(resp.Answer) != 1 { t.Fatal("expected answer from TCP fallback") } if len(warnings) != 0 { t.Errorf("unexpected warnings after successful TCP fallback: %v", warnings) } } func TestClientAlwaysTCP(t *testing.T) { var mu sync.Mutex var calls []bool c := testClient(&QueryConfig{UseTCP: true, Retries: 1}, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { mu.Lock() calls = append(calls, useTCP) mu.Unlock() return answerMsg("example.com.", "1.2.3.4"), nil }) _, _, err := c.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA) if err != nil { t.Fatalf("unexpected error: %v", err) } if len(calls) != 1 || !calls[0] { t.Errorf("expected a single TCP call, got %v", calls) } } func TestClientRetriesAreTotalAttempts(t *testing.T) { var calls int c := testClient(&QueryConfig{Retries: 3}, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { calls++ return nil, errors.New("connection refused") }) _, _, err := c.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA) if err == nil { t.Fatal("expected error after retries exhausted") } // dnsruby retry_times counts total transmissions, not extra retries. if calls != 3 { t.Errorf("expected 3 attempts, got %d", calls) } if !strings.Contains(err.Error(), "after 3 attempts") { t.Errorf("error should mention attempt count: %v", err) } } func TestClientZeroRetriesStillQueriesOnce(t *testing.T) { var calls int c := testClient(&QueryConfig{Retries: 0}, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { calls++ return nil, errors.New("boom") }) _, _, err := c.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA) if err == nil { t.Fatal("expected error") } if calls != 1 { t.Errorf("Retries=0 must be clamped to one attempt, got %d", calls) } // Regression: the old code produced "failed after 0 retries: %!w()". if strings.Contains(err.Error(), "%!w") || strings.Contains(err.Error(), "") { t.Errorf("malformed error message: %v", err) } } func TestClientRetryThenSuccess(t *testing.T) { var calls int c := testClient(&QueryConfig{Retries: 3}, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { calls++ if calls < 2 { return nil, errors.New("transient error") } return answerMsg("example.com.", "1.2.3.4"), nil }) resp, _, err := c.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA) if err != nil { t.Fatalf("unexpected error: %v", err) } if len(resp.Answer) == 0 { t.Fatal("expected answer after retry") } if calls != 2 { t.Errorf("expected 2 attempts, got %d", calls) } } func TestClientNilResponseIsError(t *testing.T) { c := testClient(&QueryConfig{Retries: 2}, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { return nil, nil }) _, _, err := c.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA) if err == nil { t.Fatal("expected error for nil responses") } } func TestClientContextCancellation(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) cancel() c := testClient(&QueryConfig{Retries: 5}, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { return nil, errors.New("error") }) _, _, err := c.Query(ctx, net.ParseIP("8.8.8.8"), "example.com", TypeA) if err == nil { t.Fatal("expected error on cancelled context") } } func TestClientNilConfigUsesDefaults(t *testing.T) { c := NewClient(nil, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { return answerMsg("example.com.", "1.2.3.4"), nil }) _, _, err := c.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA) if err != nil { t.Fatalf("unexpected error with nil config: %v", err) } } func TestClientTCPFallbackFailureRetries(t *testing.T) { var calls int c := testClient(&QueryConfig{Retries: 2, AllowTCP: true}, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { calls++ if !useTCP { m := new(dns.Msg) m.SetReply(msg) m.Truncated = true return m, nil } return nil, errors.New("tcp failed") }) _, _, err := c.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA) if err == nil { t.Fatal("expected error when TCP fallback always fails") } if calls != 4 { t.Errorf("expected 4 exchange calls (2 attempts x UDP+TCP), got %d", calls) } } func TestRetryGap(t *testing.T) { d := 2 * time.Second tests := []struct { retry int want time.Duration }{ {1, 4 * time.Second}, // dnsruby sends retry 1 at absolute 2d {2, 4 * time.Second}, // retry 2 at 4d → gap 2d {3, 8 * time.Second}, {4, 16 * time.Second}, } for _, tt := range tests { got := retryGap(d, tt.retry) if got != tt.want { t.Errorf("retryGap(%v, %d) = %v, want %v", d, tt.retry, got, tt.want) } } } func TestQueryErrorMessageMentionsQuery(t *testing.T) { c := testClient(&QueryConfig{Retries: 1}, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { return nil, errors.New("unreachable") }) _, _, err := c.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA) if err == nil { t.Fatal("expected error") } for _, part := range []string{"example.com.", "A", "8.8.8.8", "unreachable"} { if !strings.Contains(err.Error(), part) { t.Errorf("error %q should contain %q", err, part) } } }