Merge pull request 'feat: implement CachingResolver with per-server query caching' (#5) from agent/go-expert-developer/f5bcb2fe into main
CI / test (push) Failing after 3h13m21s

Reviewed-on: http://gitea.hansenits.com.au/hits/ExploreDNS/pulls/5
This commit was merged in pull request #5.
This commit is contained in:
2026-06-05 06:05:01 +00:00
2 changed files with 655 additions and 0 deletions
+198
View File
@@ -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
}
+457
View File
@@ -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")
}
}