477 lines
13 KiB
Go
477 lines
13 KiB
Go
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")
|
|
}
|
|
}
|