feat: implement CachingResolver with per-server query caching #5
@@ -0,0 +1,198 @@
|
||||
package dns
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/miekg/dns"
|
||||
)
|
||||
|
||||
type Resolver interface {
|
||||
Query(ctx context.Context, server net.IP, name string, qtype uint16, cfg *QueryConfig) (*dns.Msg, error)
|
||||
}
|
||||
|
||||
type BasicResolver struct{}
|
||||
|
||||
func NewBasicResolver() *BasicResolver {
|
||||
return &BasicResolver{}
|
||||
}
|
||||
|
||||
func (br *BasicResolver) Query(ctx context.Context, server net.IP, name string, qtype uint16, cfg *QueryConfig) (*dns.Msg, error) {
|
||||
return Query(ctx, server, name, qtype, cfg)
|
||||
}
|
||||
|
||||
type cacheKey struct {
|
||||
server string
|
||||
name string
|
||||
qtype uint16
|
||||
qclass uint16
|
||||
}
|
||||
|
||||
type cacheEntry struct {
|
||||
msg *dns.Msg
|
||||
expireAt time.Time
|
||||
}
|
||||
|
||||
func (e *cacheEntry) expired() bool {
|
||||
return time.Now().After(e.expireAt)
|
||||
}
|
||||
|
||||
type CachingResolver struct {
|
||||
inner Resolver
|
||||
mu sync.RWMutex
|
||||
cache map[cacheKey]*cacheEntry
|
||||
defaultTTL time.Duration
|
||||
}
|
||||
|
||||
func NewCachingResolver(inner Resolver, opts ...CachingResolverOption) *CachingResolver {
|
||||
if inner == nil {
|
||||
inner = NewBasicResolver()
|
||||
}
|
||||
|
||||
cfg := &cachingResolverConfig{
|
||||
defaultTTL: 5 * time.Second,
|
||||
}
|
||||
for _, opt := range opts {
|
||||
opt(cfg)
|
||||
}
|
||||
|
||||
return &CachingResolver{
|
||||
inner: inner,
|
||||
cache: make(map[cacheKey]*cacheEntry),
|
||||
defaultTTL: cfg.defaultTTL,
|
||||
}
|
||||
}
|
||||
|
||||
func (cr *CachingResolver) Query(ctx context.Context, server net.IP, name string, qtype uint16, cfg *QueryConfig) (*dns.Msg, error) {
|
||||
if cfg == nil {
|
||||
cfg = DefaultQueryConfig()
|
||||
}
|
||||
|
||||
fqdn := dns.Fqdn(name)
|
||||
key := cacheKey{
|
||||
server: server.String(),
|
||||
name: fqdn,
|
||||
qtype: qtype,
|
||||
qclass: dns.ClassINET,
|
||||
}
|
||||
|
||||
if resp, ok := cr.lookup(key); ok {
|
||||
return resp, nil
|
||||
}
|
||||
|
||||
resp, err := cr.inner.Query(ctx, server, name, qtype, cfg)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("caching resolver query: %w", err)
|
||||
}
|
||||
|
||||
cr.store(key, resp)
|
||||
return resp, nil
|
||||
}
|
||||
|
||||
func (cr *CachingResolver) lookup(key cacheKey) (*dns.Msg, bool) {
|
||||
cr.mu.RLock()
|
||||
entry, ok := cr.cache[key]
|
||||
cr.mu.RUnlock()
|
||||
if !ok {
|
||||
return nil, false
|
||||
}
|
||||
if entry.expired() {
|
||||
return nil, false
|
||||
}
|
||||
return entry.msg.Copy(), true
|
||||
}
|
||||
|
||||
func (cr *CachingResolver) store(key cacheKey, msg *dns.Msg) {
|
||||
ttl := minTTLFromMsg(msg)
|
||||
if ttl <= 0 {
|
||||
ttl = cr.defaultTTL
|
||||
}
|
||||
|
||||
cr.mu.Lock()
|
||||
cr.cache[key] = &cacheEntry{
|
||||
msg: msg.Copy(),
|
||||
expireAt: time.Now().Add(ttl),
|
||||
}
|
||||
cr.mu.Unlock()
|
||||
}
|
||||
|
||||
func (cr *CachingResolver) Len() int {
|
||||
cr.mu.RLock()
|
||||
n := len(cr.cache)
|
||||
cr.mu.RUnlock()
|
||||
return n
|
||||
}
|
||||
|
||||
func (cr *CachingResolver) Clear() {
|
||||
cr.mu.Lock()
|
||||
cr.cache = make(map[cacheKey]*cacheEntry)
|
||||
cr.mu.Unlock()
|
||||
}
|
||||
|
||||
func (cr *CachingResolver) PurgeExpired() int {
|
||||
cr.mu.Lock()
|
||||
count := 0
|
||||
for k, e := range cr.cache {
|
||||
if e.expired() {
|
||||
delete(cr.cache, k)
|
||||
count++
|
||||
}
|
||||
}
|
||||
cr.mu.Unlock()
|
||||
return count
|
||||
}
|
||||
|
||||
type CachingResolverOption func(*cachingResolverConfig)
|
||||
|
||||
type cachingResolverConfig struct {
|
||||
defaultTTL time.Duration
|
||||
}
|
||||
|
||||
func WithDefaultTTL(d time.Duration) CachingResolverOption {
|
||||
return func(c *cachingResolverConfig) {
|
||||
c.defaultTTL = d
|
||||
}
|
||||
}
|
||||
|
||||
func minTTLFromMsg(msg *dns.Msg) time.Duration {
|
||||
if msg == nil {
|
||||
return 0
|
||||
}
|
||||
|
||||
var min uint32
|
||||
found := false
|
||||
|
||||
for _, rr := range msg.Answer {
|
||||
ttl := rr.Header().Ttl
|
||||
if !found || ttl < min {
|
||||
min = ttl
|
||||
found = true
|
||||
}
|
||||
}
|
||||
for _, rr := range msg.Ns {
|
||||
ttl := rr.Header().Ttl
|
||||
if !found || ttl < min {
|
||||
min = ttl
|
||||
found = true
|
||||
}
|
||||
}
|
||||
for _, rr := range msg.Extra {
|
||||
if _, ok := rr.(*dns.OPT); ok {
|
||||
continue
|
||||
}
|
||||
ttl := rr.Header().Ttl
|
||||
if !found || ttl < min {
|
||||
min = ttl
|
||||
found = true
|
||||
}
|
||||
}
|
||||
|
||||
if !found {
|
||||
return 0
|
||||
}
|
||||
|
||||
return time.Duration(min) * time.Second
|
||||
}
|
||||
@@ -0,0 +1,457 @@
|
||||
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")
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user