From db3ad498f668afb5f72d21857ef6f6ecb6cc7161 Mon Sep 17 00:00:00 2001 From: Gary Date: Fri, 5 Jun 2026 15:58:46 +1000 Subject: [PATCH] feat: implement CachingResolver with per-server query caching Add Resolver interface, BasicResolver, and CachingResolver to internal/dns. CachingResolver wraps any Resolver with an in-memory cache keyed by (server IP, name, type, class). Cache entries respect DNS TTL from responses, falling back to a configurable default TTL. Thread-safe using sync.RWMutex. Includes Len, Clear, and PurgeExpired methods for cache management. Closes HAN-379 (Phase 2.2). Co-authored-by: multica-agent --- internal/dns/resolver.go | 198 +++++++++++++++ internal/dns/resolver_test.go | 457 ++++++++++++++++++++++++++++++++++ 2 files changed, 655 insertions(+) create mode 100644 internal/dns/resolver.go create mode 100644 internal/dns/resolver_test.go diff --git a/internal/dns/resolver.go b/internal/dns/resolver.go new file mode 100644 index 0000000..ee9ed5f --- /dev/null +++ b/internal/dns/resolver.go @@ -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 +} diff --git a/internal/dns/resolver_test.go b/internal/dns/resolver_test.go new file mode 100644 index 0000000..ee8458c --- /dev/null +++ b/internal/dns/resolver_test.go @@ -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") + } +} -- 2.54.0