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 }