package dns import ( "context" "errors" "net" "sync" "testing" "time" "github.com/miekg/dns" ) type mockResolver struct { mu sync.Mutex calls int response *dns.Msg err error } func (m *mockResolver) Query(ctx context.Context, server net.IP, name string, qtype uint16, cfg *QueryConfig) (*dns.Msg, error) { m.mu.Lock() defer m.mu.Unlock() m.calls++ if m.err != nil { return nil, m.err } if m.response != nil { return m.response.Copy(), nil } return nil, errors.New("no response configured") } func (m *mockResolver) callCount() int { m.mu.Lock() defer m.mu.Unlock() return m.calls } func makeResponse(name string, qtype uint16, ttl uint32) *dns.Msg { resp := new(dns.Msg) resp.SetReply(new(dns.Msg)) switch qtype { case dns.TypeA: resp.Answer = append(resp.Answer, &dns.A{ Hdr: dns.RR_Header{Name: dns.Fqdn(name), Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: ttl}, A: net.ParseIP("93.184.216.34"), }) case dns.TypeNS: resp.Answer = append(resp.Answer, &dns.NS{ Hdr: dns.RR_Header{Name: dns.Fqdn(name), Rrtype: dns.TypeNS, Class: dns.ClassINET, Ttl: ttl}, Ns: "ns1.example.com.", }) } return resp } func TestNewCachingResolver_NilInner(t *testing.T) { cr := NewCachingResolver(nil) if cr == nil { t.Fatal("NewCachingResolver(nil) returned nil") } if cr.inner == nil { t.Fatal("expected inner resolver to be set when nil passed") } } func TestCachingResolver_CacheHit(t *testing.T) { mock := &mockResolver{ response: makeResponse("example.com", dns.TypeA, 300), } cr := NewCachingResolver(mock) server := net.ParseIP("8.8.8.8") cfg := &QueryConfig{UDPSize: 2048, Retries: 1, Timeout: 5 * time.Second} resp1, err := cr.Query(context.Background(), server, "example.com", TypeA, cfg) if err != nil { t.Fatalf("first query: %v", err) } if mock.callCount() != 1 { t.Fatalf("expected 1 call after first query, got %d", mock.callCount()) } resp2, err := cr.Query(context.Background(), server, "example.com", TypeA, cfg) if err != nil { t.Fatalf("second query: %v", err) } if mock.callCount() != 1 { t.Fatalf("expected still 1 call after second query (cache hit), got %d", mock.callCount()) } if len(resp1.Answer) != len(resp2.Answer) { t.Errorf("cached response has different number of answers") } } func TestCachingResolver_CacheMissDifferentServer(t *testing.T) { mock := &mockResolver{ response: makeResponse("example.com", dns.TypeA, 300), } cr := NewCachingResolver(mock) cfg := &QueryConfig{UDPSize: 2048, Retries: 1, Timeout: 5 * time.Second} _, _ = cr.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA, cfg) if mock.callCount() != 1 { t.Fatalf("expected 1 call, got %d", mock.callCount()) } _, _ = cr.Query(context.Background(), net.ParseIP("1.1.1.1"), "example.com", TypeA, cfg) if mock.callCount() != 2 { t.Fatalf("expected 2 calls for different server, got %d", mock.callCount()) } } func TestCachingResolver_CacheMissDifferentName(t *testing.T) { mock := &mockResolver{ response: makeResponse("example.com", dns.TypeA, 300), } cr := NewCachingResolver(mock) server := net.ParseIP("8.8.8.8") cfg := &QueryConfig{UDPSize: 2048, Retries: 1, Timeout: 5 * time.Second} _, _ = cr.Query(context.Background(), server, "example.com", TypeA, cfg) if mock.callCount() != 1 { t.Fatalf("expected 1 call, got %d", mock.callCount()) } _, _ = cr.Query(context.Background(), server, "different.com", TypeA, cfg) if mock.callCount() != 2 { t.Fatalf("expected 2 calls for different name, got %d", mock.callCount()) } } func TestCachingResolver_CacheMissDifferentType(t *testing.T) { mock := &mockResolver{ response: makeResponse("example.com", dns.TypeA, 300), } cr := NewCachingResolver(mock) server := net.ParseIP("8.8.8.8") cfg := &QueryConfig{UDPSize: 2048, Retries: 1, Timeout: 5 * time.Second} _, _ = cr.Query(context.Background(), server, "example.com", TypeA, cfg) if mock.callCount() != 1 { t.Fatalf("expected 1 call, got %d", mock.callCount()) } _, _ = cr.Query(context.Background(), server, "example.com", TypeNS, cfg) if mock.callCount() != 2 { t.Fatalf("expected 2 calls for different qtype, got %d", mock.callCount()) } } func TestCachingResolver_ErrorNotCached(t *testing.T) { mock := &mockResolver{ err: errors.New("connection refused"), } cr := NewCachingResolver(mock) server := net.ParseIP("8.8.8.8") cfg := &QueryConfig{UDPSize: 2048, Retries: 1, Timeout: 5 * time.Second} _, err := cr.Query(context.Background(), server, "example.com", TypeA, cfg) if err == nil { t.Fatal("expected error from mock") } mock.err = nil mock.response = makeResponse("example.com", dns.TypeA, 300) _, err = cr.Query(context.Background(), server, "example.com", TypeA, cfg) if err != nil { t.Fatalf("second query after mock fixed: %v", err) } if mock.callCount() != 2 { t.Fatalf("expected 2 calls (error not cached), got %d", mock.callCount()) } } func TestCachingResolver_TTLExpiry(t *testing.T) { mock := &mockResolver{ response: makeResponse("example.com", dns.TypeA, 1), } cr := NewCachingResolver(mock, WithDefaultTTL(1*time.Second)) server := net.ParseIP("8.8.8.8") cfg := &QueryConfig{UDPSize: 2048, Retries: 1, Timeout: 5 * time.Second} _, err := cr.Query(context.Background(), server, "example.com", TypeA, cfg) if err != nil { t.Fatalf("first query: %v", err) } if mock.callCount() != 1 { t.Fatalf("expected 1 call, got %d", mock.callCount()) } time.Sleep(2 * time.Second) _, err = cr.Query(context.Background(), server, "example.com", TypeA, cfg) if err != nil { t.Fatalf("query after TTL expiry: %v", err) } if mock.callCount() != 2 { t.Fatalf("expected 2 calls after TTL expiry, got %d", mock.callCount()) } } func TestCachingResolver_DefaultTTL(t *testing.T) { resp := new(dns.Msg) resp.SetReply(new(dns.Msg)) mock := &mockResolver{response: resp} cr := NewCachingResolver(mock, WithDefaultTTL(100*time.Millisecond)) server := net.ParseIP("8.8.8.8") cfg := &QueryConfig{UDPSize: 2048, Retries: 1, Timeout: 5 * time.Second} _, err := cr.Query(context.Background(), server, "nodata.com", TypeA, cfg) if err != nil { t.Fatalf("first query: %v", err) } time.Sleep(150 * time.Millisecond) _, err = cr.Query(context.Background(), server, "nodata.com", TypeA, cfg) if err != nil { t.Fatalf("query after default TTL expiry: %v", err) } if mock.callCount() != 2 { t.Fatalf("expected 2 calls after default TTL, got %d", mock.callCount()) } } func TestCachingResolver_NilConfig(t *testing.T) { mock := &mockResolver{ response: makeResponse("example.com", dns.TypeA, 300), } cr := NewCachingResolver(mock) server := net.ParseIP("8.8.8.8") _, err := cr.Query(context.Background(), server, "example.com", TypeA, nil) if err != nil { t.Fatalf("nil config: %v", err) } } func TestCachingResolver_Len(t *testing.T) { mock := &mockResolver{ response: makeResponse("example.com", dns.TypeA, 300), } cr := NewCachingResolver(mock) if cr.Len() != 0 { t.Fatalf("expected empty cache, got %d", cr.Len()) } server := net.ParseIP("8.8.8.8") cfg := &QueryConfig{UDPSize: 2048, Retries: 1, Timeout: 5 * time.Second} _, _ = cr.Query(context.Background(), server, "example.com", TypeA, cfg) if cr.Len() != 1 { t.Fatalf("expected cache len 1, got %d", cr.Len()) } _, _ = cr.Query(context.Background(), server, "other.com", TypeA, cfg) if cr.Len() != 2 { t.Fatalf("expected cache len 2, got %d", cr.Len()) } } func TestCachingResolver_Clear(t *testing.T) { mock := &mockResolver{ response: makeResponse("example.com", dns.TypeA, 300), } cr := NewCachingResolver(mock) server := net.ParseIP("8.8.8.8") cfg := &QueryConfig{UDPSize: 2048, Retries: 1, Timeout: 5 * time.Second} _, _ = cr.Query(context.Background(), server, "example.com", TypeA, cfg) _, _ = cr.Query(context.Background(), server, "other.com", TypeA, cfg) if cr.Len() != 2 { t.Fatalf("expected cache len 2, got %d", cr.Len()) } cr.Clear() if cr.Len() != 0 { t.Fatalf("expected cache len 0 after clear, got %d", cr.Len()) } _, err := cr.Query(context.Background(), server, "example.com", TypeA, cfg) if err != nil { t.Fatalf("query after clear: %v", err) } if mock.callCount() != 3 { t.Fatalf("expected 3 calls (2 before clear + 1 after clear), got %d", mock.callCount()) } } func TestCachingResolver_PurgeExpired(t *testing.T) { resp := makeResponse("example.com", dns.TypeA, 0) for _, rr := range resp.Answer { rr.Header().Ttl = 0 } mock := &mockResolver{response: resp} cr := NewCachingResolver(mock, WithDefaultTTL(50*time.Millisecond)) server := net.ParseIP("8.8.8.8") cfg := &QueryConfig{UDPSize: 2048, Retries: 1, Timeout: 5 * time.Second} _, _ = cr.Query(context.Background(), server, "example.com", TypeA, cfg) if cr.Len() != 1 { t.Fatalf("expected cache len 1, got %d", cr.Len()) } time.Sleep(100 * time.Millisecond) purged := cr.PurgeExpired() if purged != 1 { t.Fatalf("expected 1 purged entry, got %d", purged) } if cr.Len() != 0 { t.Fatalf("expected cache len 0 after purge, got %d", cr.Len()) } } func TestCachingResolver_PurgeExpiredNoneExpired(t *testing.T) { mock := &mockResolver{ response: makeResponse("example.com", dns.TypeA, 300), } cr := NewCachingResolver(mock) server := net.ParseIP("8.8.8.8") cfg := &QueryConfig{UDPSize: 2048, Retries: 1, Timeout: 5 * time.Second} _, _ = cr.Query(context.Background(), server, "example.com", TypeA, cfg) purged := cr.PurgeExpired() if purged != 0 { t.Fatalf("expected 0 purged entries, got %d", purged) } if cr.Len() != 1 { t.Fatalf("expected cache len 1, got %d", cr.Len()) } } func TestCachingResolver_ResponseCopy(t *testing.T) { mock := &mockResolver{ response: makeResponse("example.com", dns.TypeA, 300), } cr := NewCachingResolver(mock) server := net.ParseIP("8.8.8.8") cfg := &QueryConfig{UDPSize: 2048, Retries: 1, Timeout: 5 * time.Second} resp1, _ := cr.Query(context.Background(), server, "example.com", TypeA, cfg) resp2, _ := cr.Query(context.Background(), server, "example.com", TypeA, cfg) if resp1 == resp2 { t.Fatal("cache should return a copy, not the same pointer") } } func TestCachingResolver_MinTTLFromMsg(t *testing.T) { tests := []struct { name string msg *dns.Msg expected time.Duration }{ { name: "nil message", msg: nil, expected: 0, }, { name: "empty response", msg: new(dns.Msg), expected: 0, }, { name: "single answer with low TTL", msg: func() *dns.Msg { m := new(dns.Msg) m.Answer = append(m.Answer, &dns.A{ Hdr: dns.RR_Header{Ttl: 60}, }) return m }(), expected: 60 * time.Second, }, { name: "multiple records with varying TTLs", msg: func() *dns.Msg { m := new(dns.Msg) m.Answer = append(m.Answer, &dns.A{ Hdr: dns.RR_Header{Ttl: 300}, }) m.Ns = append(m.Ns, &dns.NS{ Hdr: dns.RR_Header{Ttl: 120}, }) return m }(), expected: 120 * time.Second, }, { name: "OPT record excluded from TTL calculation", msg: func() *dns.Msg { m := new(dns.Msg) m.Answer = append(m.Answer, &dns.A{ Hdr: dns.RR_Header{Ttl: 300}, }) m.Extra = append(m.Extra, &dns.OPT{ Hdr: dns.RR_Header{Ttl: 0}, }) return m }(), expected: 300 * time.Second, }, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { got := minTTLFromMsg(tt.msg) if got != tt.expected { t.Errorf("minTTLFromMsg() = %v, want %v", got, tt.expected) } }) } } func TestCachingResolver_ConcurrentAccess(t *testing.T) { mock := &mockResolver{ response: makeResponse("example.com", dns.TypeA, 300), } cr := NewCachingResolver(mock) server := net.ParseIP("8.8.8.8") cfg := &QueryConfig{UDPSize: 2048, Retries: 1, Timeout: 5 * time.Second} var wg sync.WaitGroup for i := 0; i < 100; i++ { wg.Add(1) go func() { defer wg.Done() _, err := cr.Query(context.Background(), server, "example.com", TypeA, cfg) if err != nil { t.Errorf("concurrent query failed: %v", err) } }() } wg.Wait() if mock.callCount() < 1 { t.Fatalf("expected at least 1 call to mock, got %d", mock.callCount()) } } func TestBasicResolver(t *testing.T) { br := NewBasicResolver() if br == nil { t.Fatal("NewBasicResolver returned nil") } if _, ok := interface{}(br).(Resolver); !ok { t.Fatal("BasicResolver does not implement Resolver interface") } } // TestBasicResolverQueryIntegration calls Query via the BasicResolver against // the local system resolver. Skipped when no local resolver is reachable. func TestBasicResolverQueryIntegration(t *testing.T) { r := NewBasicResolver() server := net.ParseIP("127.0.0.1") ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() // Cover BasicResolver.Query; skip if 127.0.0.1:53 is not available. msg, err := r.Query(ctx, server, ".", TypeNS, nil) if err != nil { t.Logf("skipping (local resolver unavailable): %v", err) t.Skip() } if msg == nil { t.Fatal("expected non-nil response from BasicResolver.Query") } }