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>
199 lines
3.5 KiB
Go
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
|
|
}
|