feat: implement CachingResolver with per-server query caching
CI / test (pull_request) Failing after 3h13m39s
CI / test (pull_request) Failing after 3h13m39s
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>
This commit is contained in:
@@ -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
|
||||
}
|
||||
Reference in New Issue
Block a user