Files
ExploreDNS/internal/dns/resolver.go
Garyandmultica-agent db3ad498f6
CI / test (pull_request) Failing after 3h13m39s
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 <github@multica.ai>
2026-06-05 15:58:46 +10:00

199 lines
3.5 KiB
Go

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
}