package dns import ( "context" "fmt" "net" "sync" "time" "github.com/miekg/dns" ) // Resolver is the interface for sending DNS queries. // Implementations may cache, mock, or delegate to the real network. type Resolver interface { Query(ctx context.Context, server net.IP, name string, qtype uint16, cfg *QueryConfig) (*dns.Msg, error) } // BasicResolver implements Resolver by issuing real UDP/TCP DNS queries. type BasicResolver struct{} // NewBasicResolver returns a BasicResolver ready for use. 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) } // CachingResolver wraps another Resolver and caches successful responses. // Responses are evicted based on the minimum DNS TTL in the message, with a // configurable fallback default TTL for zero-TTL responses. // All operations are safe for concurrent use. type CachingResolver struct { inner Resolver mu sync.RWMutex cache map[cacheKey]*cacheEntry defaultTTL time.Duration } // NewCachingResolver creates a CachingResolver wrapping inner. // When inner is nil, a BasicResolver is used. // Optional opts may be supplied to customise the default TTL (see WithDefaultTTL). 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() } // Len returns the current number of entries in the cache, including expired ones. func (cr *CachingResolver) Len() int { cr.mu.RLock() n := len(cr.cache) cr.mu.RUnlock() return n } // Clear removes all entries from the cache. func (cr *CachingResolver) Clear() { cr.mu.Lock() cr.cache = make(map[cacheKey]*cacheEntry) cr.mu.Unlock() } // PurgeExpired removes all expired entries from the cache and returns the count removed. 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 } // CachingResolverOption is a functional option for NewCachingResolver. type CachingResolverOption func(*cachingResolverConfig) type cachingResolverConfig struct { defaultTTL time.Duration } // WithDefaultTTL sets the TTL used for cache entries when the DNS response // contains no TTL information (or when TTL is zero). 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 }