fix(cache): harden providers and default cache handling

Make cache-write failures non-fatal in GetOrSet/Remember, replace the
unsynchronised default cache with an atomic pointer, fix tag-index leaks
and write-lock-on-read in the memory provider, add a janitor, closed
state and byte copies, hash/namespace memcache keys with CAS tag index,
gate Clear behind AllowFlush, return ErrNotFound without the key, and
allowlist Redis stats. Mark audit status.
This commit is contained in:
Hein
2026-09-30 13:13:06 +02:00
parent 652621a70e
commit 9533c3a0ed
9 changed files with 648 additions and 262 deletions
+29
View File
@@ -64,6 +64,35 @@ Every error is either returned to the caller or discarded with `_ =`.
| 24 | **Low** | correctness | No defensive copy of `[]byte` on `Set`/`Get` in the memory provider |
| 25 | **Low** | hygiene | `example_usage.go` ships `log.Fatal` calls in a library package |
## Resolution status (2026-09-30)
- **#1** — Fixed: cache-write failure is logged and ignored in `GetOrSet`/`Remember`
- **#2** — Not fixed here: needs `pkg/security` changes (tag-based revocation); memcache `DeleteByPattern` still errors (now documented)
- **#3** — Fixed: `atomic.Pointer` + CAS lazy init; displaced provider closed on `Initialize`/`Use*` (not on `SetDefaultCache`, where the caller owns it)
- **#4** — Fixed: single `removeLocked` used by every removal path; `Clear` resets the index
- **#5** — Fixed: `Get` runs under `RLock`; access counters are atomics
- **#6** — Deferred (pattern language is an API decision). Memory now compiles the regexp before locking; Redis `SCAN COUNT` is 500
- **#7** — Fixed: `ErrNotFound` sentinel, no key in the error. `pkg/security` still keys on the raw token (out of scope)
- **#8** — Deferred (needs singleflight dependency)
- **#9** — Fixed in the memcache provider: keys are namespaced (`k:`) and hashed when over 200 bytes or illegal
- **#10** — Fixed: re-check under the write lock before deleting
- **#11** — Fixed: `Close` sets a closed flag; writes return `ErrClosed`
- **#12** — Fixed: CAS retry loop, bounded tag list (5000), errors returned, 30-day expiry rule, user keys namespaced. Failed tag indexing rolls back the value
- **#13** — Fixed: `Clear` on Redis/Memcache returns `ErrFlushNotAllowed` unless `AllowFlush` is set (behaviour change)
- **#14** — Fixed: `Close` calls `client.Close()`
- **#15** — Deferred (new TLS config fields)
- **#16** — Documented only: warning on `Remember`; signature unchanged
- **#17** — Partly fixed: Redis/Memcache `Get` now log backend errors; the `Provider` interface is unchanged so it is still reported as a miss
- **#18** — Not fixed: still an O(n) scan (comment fixed in `Stats`)
- **#19** — Fixed: janitor goroutine, `Options.CleanupInterval` (default 1m), stopped by `Close`
- **#20** — Fixed: only allowlisted counters are exposed; `Hits`/`Misses` populated
- **#21** — Partly fixed: memcache methods check `ctx.Err()` first; in-flight calls are still bounded only by `Timeout`
- **#22** — Fixed: `MaxSize` 0 means 10000; negative means unbounded
- **#23** — Fixed: constructors copy config and `Options`
- **#24** — Fixed: memory provider copies on `Set` and `Get`
- **#25** — Fixed: `example_usage.go` excluded from the build with `//go:build ignore`
- Tests: `pkg/cache/hardening_test.go` (run with `-race`).
---
### 1. Critical — a cache-write failure is an authentication failure
+30 -17
View File
@@ -3,23 +3,29 @@ package cache
import (
"context"
"fmt"
"sync/atomic"
"time"
)
var (
defaultCache *Cache
)
var defaultCache atomic.Pointer[Cache]
// swapOwned installs c as the default and closes the displaced cache, which this
// package created and therefore owns.
func swapOwned(c *Cache) {
if old := defaultCache.Swap(c); old != nil && old != c {
_ = old.Close() // best-effort: the displaced provider is being discarded
}
}
// Initialize initializes the cache with a provider.
// If not called, the package will use an in-memory provider by default.
func Initialize(provider Provider) {
defaultCache = NewCache(provider)
swapOwned(NewCache(provider))
}
// UseMemory configures the cache to use in-memory storage.
func UseMemory(opts *Options) error {
provider := NewMemoryProvider(opts)
defaultCache = NewCache(provider)
swapOwned(NewCache(NewMemoryProvider(opts)))
return nil
}
@@ -29,7 +35,7 @@ func UseRedis(config *RedisConfig) error {
if err != nil {
return fmt.Errorf("failed to initialize Redis provider: %w", err)
}
defaultCache = NewCache(provider)
swapOwned(NewCache(provider))
return nil
}
@@ -39,26 +45,33 @@ func UseMemcache(config *MemcacheConfig) error {
if err != nil {
return fmt.Errorf("failed to initialize Memcache provider: %w", err)
}
defaultCache = NewCache(provider)
swapOwned(NewCache(provider))
return nil
}
// GetDefaultCache returns the default cache instance.
// Initializes with in-memory provider if not already initialized.
// Safe for concurrent use.
func GetDefaultCache() *Cache {
if defaultCache == nil {
_ = UseMemory(&Options{
DefaultTTL: 5 * time.Minute,
MaxSize: 10000,
})
if c := defaultCache.Load(); c != nil {
return c
}
return defaultCache
fresh := NewCache(NewMemoryProvider(&Options{
DefaultTTL: 5 * time.Minute,
MaxSize: 10000,
}))
if defaultCache.CompareAndSwap(nil, fresh) {
return fresh
}
_ = fresh.Close() // lost the race; discard our provider
return defaultCache.Load()
}
// SetDefaultCache sets a custom cache instance as the default cache.
// This is useful for testing or when you want to use a pre-configured cache instance.
// The caller keeps ownership of both the new and the displaced cache; neither is closed.
func SetDefaultCache(cache *Cache) {
defaultCache = cache
defaultCache.Store(cache)
}
// GetStats returns cache statistics.
@@ -69,8 +82,8 @@ func GetStats(ctx context.Context) (*CacheStats, error) {
// Close closes the cache and releases resources.
func Close() error {
if defaultCache != nil {
return defaultCache.Close()
if c := defaultCache.Load(); c != nil {
return c.Close()
}
return nil
}
+19 -6
View File
@@ -3,10 +3,17 @@ package cache
import (
"context"
"encoding/json"
"errors"
"fmt"
"time"
"github.com/bitechdev/ResolveSpec/pkg/logger"
)
// ErrNotFound is returned when a key is not in the cache. The key is deliberately
// not included in the error: keys may embed credentials (e.g. session tokens).
var ErrNotFound = errors.New("cache: key not found")
// Cache is the main cache manager that wraps a Provider.
type Cache struct {
provider Provider
@@ -23,7 +30,7 @@ func NewCache(provider Provider) *Cache {
func (c *Cache) Get(ctx context.Context, key string, dest interface{}) error {
data, exists := c.provider.Get(ctx, key)
if !exists {
return fmt.Errorf("key not found: %s", key)
return ErrNotFound
}
if err := json.Unmarshal(data, dest); err != nil {
@@ -37,7 +44,7 @@ func (c *Cache) Get(ctx context.Context, key string, dest interface{}) error {
func (c *Cache) GetBytes(ctx context.Context, key string) ([]byte, error) {
data, exists := c.provider.Get(ctx, key)
if !exists {
return nil, fmt.Errorf("key not found: %s", key)
return nil, ErrNotFound
}
return data, nil
}
@@ -122,9 +129,10 @@ func (c *Cache) GetOrSet(ctx context.Context, key string, dest interface{}, ttl
return fmt.Errorf("loader failed: %w", err)
}
// Store in cache
// Store in cache. A cache-write failure must not fail the call: the authoritative
// value has already been loaded, and the cache is only an optimisation.
if err := c.Set(ctx, key, value, ttl); err != nil {
return fmt.Errorf("failed to cache value: %w", err)
logger.Warn("cache: failed to store loaded value, continuing uncached: %v", err)
}
// Populate dest with the loaded value
@@ -142,6 +150,11 @@ func (c *Cache) GetOrSet(ctx context.Context, key string, dest interface{}, ttl
// Remember is a convenience function that caches the result of a function call.
// It's similar to GetOrSet but returns the value directly.
//
// WARNING: the returned type differs between a hit and a miss. On a hit the value is
// generic decoded JSON (map[string]interface{}, []interface{}, float64, string, ...);
// on a miss it is exactly what loader returned. Do not type-assert the result to a
// concrete type; prefer GetOrSet, which decodes into a typed destination.
func (c *Cache) Remember(ctx context.Context, key string, ttl time.Duration, loader func() (interface{}, error)) (interface{}, error) {
// Try to get from cache first as bytes
data, err := c.GetBytes(ctx, key)
@@ -158,9 +171,9 @@ func (c *Cache) Remember(ctx context.Context, key string, ttl time.Duration, loa
return nil, fmt.Errorf("loader failed: %w", err)
}
// Store in cache
// Cache-write failures are non-fatal (see GetOrSet)
if err := c.Set(ctx, key, value, ttl); err != nil {
return nil, fmt.Errorf("failed to cache value: %w", err)
logger.Warn("cache: failed to store loaded value, continuing uncached: %v", err)
}
return value, nil
+4
View File
@@ -1,3 +1,7 @@
//go:build ignore
// Examples are excluded from the build: they call log.Fatal and are not part of the API.
package cache
import (
+154
View File
@@ -0,0 +1,154 @@
package cache
import (
"context"
"errors"
"fmt"
"strings"
"sync"
"testing"
"time"
)
type failSetProvider struct{ *MemoryProvider }
func (f failSetProvider) Set(context.Context, string, []byte, time.Duration) error {
return errors.New("backend down")
}
func TestGetOrSetSurvivesCacheWriteFailure(t *testing.T) {
c := NewCache(failSetProvider{NewMemoryProvider(nil)})
var out string
err := c.GetOrSet(context.Background(), "k", &out, time.Minute, func() (interface{}, error) { return "v", nil })
if err != nil || out != "v" {
t.Fatalf("got %q, %v", out, err)
}
}
func TestNotFoundDoesNotLeakKey(t *testing.T) {
c := NewCache(NewMemoryProvider(nil))
err := c.Get(context.Background(), "auth:session:SECRET", new(string))
if !errors.Is(err, ErrNotFound) || strings.Contains(err.Error(), "SECRET") {
t.Fatalf("unexpected error %v", err)
}
}
func TestMemoryTagIndexCleanedOnAllRemovals(t *testing.T) {
ctx := context.Background()
m := NewMemoryProvider(&Options{MaxSize: 2})
defer m.Close()
for i := 0; i < 50; i++ {
_ = m.SetWithTags(ctx, fmt.Sprintf("k%d", i), []byte("x"), time.Minute, []string{"t"})
}
m.mu.RLock()
n := len(m.tagToKeys["t"])
m.mu.RUnlock()
if n > 2 {
t.Fatalf("tag index leaked: %d members with MaxSize 2", n)
}
_ = m.Clear(ctx)
if len(m.tagToKeys) != 0 {
t.Fatal("Clear did not reset tag index")
}
}
func TestMemoryClosedNoPanic(t *testing.T) {
m := NewMemoryProvider(nil)
_ = m.Close()
if err := m.Set(context.Background(), "k", []byte("v"), 0); !errors.Is(err, ErrClosed) {
t.Fatalf("got %v", err)
}
if _, ok := m.Get(context.Background(), "k"); ok {
t.Fatal("hit after close")
}
}
func TestMemoryCopiesAndDefaults(t *testing.T) {
ctx := context.Background()
opts := &Options{}
m := NewMemoryProvider(opts)
defer m.Close()
if opts.MaxSize != 0 || m.options.MaxSize != defaultMemoryMaxSize {
t.Fatal("options not copied/defaulted")
}
buf := []byte("abc")
_ = m.Set(ctx, "k", buf, time.Minute)
buf[0] = 'X'
got, _ := m.Get(ctx, "k")
if string(got) != "abc" {
t.Fatalf("stored slice aliased caller: %q", got)
}
got[0] = 'Y'
if again, _ := m.Get(ctx, "k"); string(again) != "abc" {
t.Fatal("returned slice aliases stored value")
}
}
func TestMemoryJanitorRemovesExpired(t *testing.T) {
m := NewMemoryProvider(&Options{CleanupInterval: 10 * time.Millisecond})
defer m.Close()
_ = m.Set(context.Background(), "k", []byte("v"), 5*time.Millisecond)
time.Sleep(100 * time.Millisecond)
m.mu.RLock()
n := len(m.items)
m.mu.RUnlock()
if n != 0 {
t.Fatalf("expired item still stored: %d", n)
}
}
func TestMemoryConcurrent(t *testing.T) {
ctx := context.Background()
m := NewMemoryProvider(&Options{MaxSize: 50})
defer m.Close()
var wg sync.WaitGroup
for g := 0; g < 8; g++ {
wg.Add(1)
go func(g int) {
defer wg.Done()
for i := 0; i < 300; i++ {
k := fmt.Sprintf("k%d", i%80)
_ = m.SetWithTags(ctx, k, []byte("v"), time.Millisecond, []string{"t"})
m.Get(ctx, k)
if i%50 == 0 {
_ = m.DeleteByTag(ctx, "t")
}
}
}(g)
}
wg.Wait()
}
func TestGetDefaultCacheConcurrent(t *testing.T) {
SetDefaultCache(nil)
var wg sync.WaitGroup
res := make([]*Cache, 16)
for i := range res {
wg.Add(1)
go func(i int) { defer wg.Done(); res[i] = GetDefaultCache() }(i)
}
wg.Wait()
for _, c := range res {
if c != res[0] {
t.Fatal("different default caches returned")
}
}
}
func TestMemcacheKeyAndExpiry(t *testing.T) {
if k := memcacheKey(strings.Repeat("a", 300)); len(k) > 250 || !legalMemcacheKey(k) {
t.Fatalf("bad key %q", k)
}
if k := memcacheKey("has space"); !legalMemcacheKey(k) {
t.Fatal("illegal key not normalised")
}
if memcacheKey("a") != "k:a" {
t.Fatal("unexpected prefix")
}
if got := memcacheExpiry(30*24*time.Hour + time.Hour); got < int32(time.Now().Unix()) {
t.Fatalf("expected absolute timestamp, got %d", got)
}
if memcacheExpiry(-time.Second) != 0 || memcacheExpiry(time.Minute) != 60 {
t.Fatal("bad relative expiry")
}
}
+10
View File
@@ -2,9 +2,14 @@ package cache
import (
"context"
"errors"
"time"
)
// ErrFlushNotAllowed is returned by Clear on shared-server providers (Redis, Memcache)
// unless AllowFlush is set in their config, because Clear flushes the whole server/DB.
var ErrFlushNotAllowed = errors.New("cache: Clear flushes the entire server; set AllowFlush in the provider config to permit it")
// Provider defines the interface that all cache providers must implement.
type Provider interface {
// Get retrieves a value from the cache by key.
@@ -58,8 +63,13 @@ type Options struct {
DefaultTTL time.Duration
// MaxSize is the maximum number of items (for in-memory provider).
// 0 selects the default (10000); a negative value means unbounded.
MaxSize int
// CleanupInterval is how often the in-memory provider removes expired items
// (default: 1 minute).
CleanupInterval time.Duration
// EvictionPolicy determines how items are evicted (LRU, LFU, etc).
EvictionPolicy string
}
+237 -119
View File
@@ -2,17 +2,33 @@ package cache
import (
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"time"
"github.com/bitechdev/ResolveSpec/pkg/logger"
"github.com/bradfitz/gomemcache/memcache"
)
const (
// memcacheMaxRelativeTTL is the largest expiry memcached treats as relative seconds;
// anything larger is interpreted as an absolute Unix timestamp.
memcacheMaxRelativeTTL = 30 * 24 * 60 * 60
// memcacheMaxTagKeys bounds the per-tag key list so it stays under memcached's 1MB item limit.
memcacheMaxTagKeys = 5000
memcacheCASRetries = 5
)
// MemcacheProvider is a Memcache implementation of the Provider interface.
type MemcacheProvider struct {
client *memcache.Client
options *Options
client *memcache.Client
options *Options
allowFlush bool
}
// MemcacheConfig contains Memcache-specific configuration.
@@ -28,37 +44,46 @@ type MemcacheConfig struct {
// Options contains general cache options
Options *Options
// AllowFlush permits Clear() to run flush_all, which wipes every key on every
// configured server (including data not owned by this cache). Off by default.
AllowFlush bool
}
// NewMemcacheProvider creates a new Memcache cache provider.
func NewMemcacheProvider(config *MemcacheConfig) (*MemcacheProvider, error) {
if config == nil {
config = &MemcacheConfig{
Servers: []string{"localhost:11211"},
}
// Work on a copy so the caller's struct is not mutated
var cfg MemcacheConfig
if config != nil {
cfg = *config
cfg.Servers = append([]string(nil), config.Servers...)
}
if cfg.Options != nil {
o := *cfg.Options
cfg.Options = &o
}
if len(config.Servers) == 0 {
config.Servers = []string{"localhost:11211"}
if len(cfg.Servers) == 0 {
cfg.Servers = []string{"localhost:11211"}
}
if config.MaxIdleConns == 0 {
config.MaxIdleConns = 2
if cfg.MaxIdleConns == 0 {
cfg.MaxIdleConns = 2
}
if config.Timeout == 0 {
config.Timeout = 1 * time.Second
if cfg.Timeout == 0 {
cfg.Timeout = 1 * time.Second
}
if config.Options == nil {
config.Options = &Options{
if cfg.Options == nil {
cfg.Options = &Options{
DefaultTTL: 5 * time.Minute,
}
}
client := memcache.New(config.Servers...)
client.MaxIdleConns = config.MaxIdleConns
client.Timeout = config.Timeout
client := memcache.New(cfg.Servers...)
client.MaxIdleConns = cfg.MaxIdleConns
client.Timeout = cfg.Timeout
// Test connection
if err := client.Ping(); err != nil {
@@ -66,18 +91,72 @@ func NewMemcacheProvider(config *MemcacheConfig) (*MemcacheProvider, error) {
}
return &MemcacheProvider{
client: client,
options: config.Options,
client: client,
options: cfg.Options,
allowFlush: cfg.AllowFlush,
}, nil
}
// memcacheKey maps a caller key to a legal memcached key. User keys live under the "k:"
// prefix, so they can never collide with the "cache:tag:" / "cache:tags:" index keys.
// Keys that are too long or contain illegal bytes (whitespace/control characters)
// are replaced with their SHA-256.
func memcacheKey(key string) string {
k := "k:" + key
if len(k) > 200 || !legalMemcacheKey(k) {
sum := sha256.Sum256([]byte(key))
return "k:h:" + hex.EncodeToString(sum[:])
}
return k
}
func legalMemcacheKey(key string) bool {
for i := 0; i < len(key); i++ {
if key[i] <= ' ' || key[i] == 0x7f {
return false
}
}
return true
}
// memcacheTagKey maps a tag to its index key, hashing if it is not a legal key.
func memcacheTagKey(prefix, name string) string {
k := prefix + name
if len(k) > 200 || !legalMemcacheKey(k) {
sum := sha256.Sum256([]byte(name))
return prefix + "h:" + hex.EncodeToString(sum[:])
}
return k
}
// memcacheExpiry converts a TTL into a memcached expiry value, switching to an
// absolute Unix timestamp above 30 days as the protocol requires.
func memcacheExpiry(ttl time.Duration) int32 {
if ttl <= 0 {
return 0 // never expires
}
secs := int64(ttl.Seconds())
if secs > memcacheMaxRelativeTTL {
return int32(time.Now().Add(ttl).Unix())
}
if secs == 0 {
secs = 1
}
return int32(secs)
}
// Get retrieves a value from the cache by key.
func (m *MemcacheProvider) Get(ctx context.Context, key string) ([]byte, bool) {
item, err := m.client.Get(key)
if err == memcache.ErrCacheMiss {
if ctx.Err() != nil {
return nil, false
}
item, err := m.client.Get(memcacheKey(key))
if errors.Is(err, memcache.ErrCacheMiss) {
return nil, false
}
if err != nil {
// Reported as a miss (the Provider interface cannot express errors), but not silently
logger.Warn("cache: memcache GET failed: %v", err)
return nil, false
}
return item.Value, true
@@ -85,130 +164,158 @@ func (m *MemcacheProvider) Get(ctx context.Context, key string) ([]byte, bool) {
// Set stores a value in the cache with the specified TTL.
func (m *MemcacheProvider) Set(ctx context.Context, key string, value []byte, ttl time.Duration) error {
if err := ctx.Err(); err != nil {
return err
}
if ttl == 0 {
ttl = m.options.DefaultTTL
}
item := &memcache.Item{
Key: key,
return m.client.Set(&memcache.Item{
Key: memcacheKey(key),
Value: value,
Expiration: int32(ttl.Seconds()),
}
return m.client.Set(item)
Expiration: memcacheExpiry(ttl),
})
}
// SetWithTags stores a value in the cache with the specified TTL and tags.
// Note: Tag support in Memcache is limited and less efficient than Redis.
// Note: Tag support in Memcache is limited and less efficient than Redis. The
// tag index is updated with compare-and-swap; if it cannot be updated the value is
// removed again and an error is returned, so an untracked entry is never left behind.
func (m *MemcacheProvider) SetWithTags(ctx context.Context, key string, value []byte, ttl time.Duration, tags []string) error {
if err := ctx.Err(); err != nil {
return err
}
if ttl == 0 {
ttl = m.options.DefaultTTL
}
expiration := int32(ttl.Seconds())
expiration := memcacheExpiry(ttl)
mkey := memcacheKey(key)
// Set the main value
item := &memcache.Item{
Key: key,
Value: value,
Expiration: expiration,
if err := m.client.Set(&memcache.Item{Key: mkey, Value: value, Expiration: expiration}); err != nil {
return err
}
if err := m.client.Set(item); err != nil {
if len(tags) == 0 {
return nil
}
fail := func(err error) error {
_ = m.client.Delete(mkey) // best-effort rollback; the original error is what matters
return err
}
// Store tags for this key
if len(tags) > 0 {
tagsData, err := json.Marshal(tags)
if err != nil {
return fmt.Errorf("failed to marshal tags: %w", err)
}
tagsData, err := json.Marshal(tags)
if err != nil {
return fail(fmt.Errorf("failed to marshal tags: %w", err))
}
if err := m.client.Set(&memcache.Item{
Key: memcacheTagKey("cache:tags:", key),
Value: tagsData,
Expiration: expiration,
}); err != nil {
return fail(err)
}
tagsItem := &memcache.Item{
Key: fmt.Sprintf("cache:tags:%s", key),
Value: tagsData,
Expiration: expiration,
}
if err := m.client.Set(tagsItem); err != nil {
return err
}
// Add key to each tag's key list
for _, tag := range tags {
tagKey := fmt.Sprintf("cache:tag:%s", tag)
// Get existing keys for this tag
var keys []string
if item, err := m.client.Get(tagKey); err == nil {
_ = json.Unmarshal(item.Value, &keys)
}
// Add current key if not already present
found := false
// Tag lists live longer than the entries they index
tagExpiry := memcacheExpiry(ttl + time.Hour)
for _, tag := range tags {
if err := m.updateTagKeys(memcacheTagKey("cache:tag:", tag), tagExpiry, func(keys []string) ([]string, error) {
for _, k := range keys {
if k == key {
found = true
break
return keys, nil
}
}
if !found {
keys = append(keys, key)
if len(keys) >= memcacheMaxTagKeys {
return nil, fmt.Errorf("tag index for %q is full (%d keys)", tag, memcacheMaxTagKeys)
}
// Store updated key list
keysData, err := json.Marshal(keys)
if err != nil {
continue
}
tagItem := &memcache.Item{
Key: tagKey,
Value: keysData,
Expiration: expiration + 3600, // Give tag lists longer TTL
}
_ = m.client.Set(tagItem)
return append(keys, key), nil
}); err != nil {
return fail(err)
}
}
return nil
}
// updateTagKeys applies fn to a tag's key list using compare-and-swap.
func (m *MemcacheProvider) updateTagKeys(tagKey string, expiry int32, fn func([]string) ([]string, error)) error {
for attempt := 0; attempt < memcacheCASRetries; attempt++ {
item, err := m.client.Get(tagKey)
var keys []string
switch {
case errors.Is(err, memcache.ErrCacheMiss):
item = nil
case err != nil:
return err
default:
if err := json.Unmarshal(item.Value, &keys); err != nil {
return fmt.Errorf("failed to unmarshal tag keys: %w", err)
}
}
keys, err = fn(keys)
if err != nil {
return err
}
data, err := json.Marshal(keys)
if err != nil {
return err
}
if item == nil {
err = m.client.Add(&memcache.Item{Key: tagKey, Value: data, Expiration: expiry})
if errors.Is(err, memcache.ErrNotStored) {
continue // someone created it first; retry
}
return err
}
item.Value = data
item.Expiration = expiry
err = m.client.CompareAndSwap(item)
if errors.Is(err, memcache.ErrCASConflict) || errors.Is(err, memcache.ErrNotStored) || errors.Is(err, memcache.ErrCacheMiss) {
continue
}
return err
}
return fmt.Errorf("tag index %q: too much contention, giving up", tagKey)
}
// Delete removes a key from the cache.
func (m *MemcacheProvider) Delete(ctx context.Context, key string) error {
if err := ctx.Err(); err != nil {
return err
}
mkey := memcacheKey(key)
// Get tags for this key
tagsKey := fmt.Sprintf("cache:tags:%s", key)
tagsKey := memcacheTagKey("cache:tags:", key)
if item, err := m.client.Get(tagsKey); err == nil {
var tags []string
if err := json.Unmarshal(item.Value, &tags); err == nil {
// Remove key from each tag's key list
for _, tag := range tags {
tagKey := fmt.Sprintf("cache:tag:%s", tag)
if tagItem, err := m.client.Get(tagKey); err == nil {
var keys []string
if err := json.Unmarshal(tagItem.Value, &keys); err == nil {
// Remove current key from the list
newKeys := make([]string, 0, len(keys))
for _, k := range keys {
if k != key {
newKeys = append(newKeys, k)
}
}
// Update the tag's key list
if keysData, err := json.Marshal(newKeys); err == nil {
tagItem.Value = keysData
_ = m.client.Set(tagItem)
err := m.updateTagKeys(memcacheTagKey("cache:tag:", tag), memcacheExpiry(m.options.DefaultTTL+time.Hour), func(keys []string) ([]string, error) {
out := make([]string, 0, len(keys))
for _, k := range keys {
if k != key {
out = append(out, k)
}
}
return out, nil
})
if err != nil {
logger.Warn("cache: failed to update memcache tag index on delete: %v", err)
}
}
}
// Delete the tags key
_ = m.client.Delete(tagsKey)
if err := m.client.Delete(tagsKey); err != nil && !errors.Is(err, memcache.ErrCacheMiss) {
logger.Warn("cache: failed to delete memcache tags key: %v", err)
}
}
// Delete the actual key
err := m.client.Delete(key)
if err == memcache.ErrCacheMiss {
err := m.client.Delete(mkey)
if errors.Is(err, memcache.ErrCacheMiss) {
return nil
}
return err
@@ -216,11 +323,13 @@ func (m *MemcacheProvider) Delete(ctx context.Context, key string) error {
// DeleteByTag removes all keys associated with the given tag.
func (m *MemcacheProvider) DeleteByTag(ctx context.Context, tag string) error {
tagKey := fmt.Sprintf("cache:tag:%s", tag)
if err := ctx.Err(); err != nil {
return err
}
tagKey := memcacheTagKey("cache:tag:", tag)
// Get all keys associated with this tag
item, err := m.client.Get(tagKey)
if err == memcache.ErrCacheMiss {
if errors.Is(err, memcache.ErrCacheMiss) {
return nil
}
if err != nil {
@@ -232,42 +341,51 @@ func (m *MemcacheProvider) DeleteByTag(ctx context.Context, tag string) error {
return fmt.Errorf("failed to unmarshal tag keys: %w", err)
}
// Delete all keys
var firstErr error
note := func(err error) {
if err != nil && !errors.Is(err, memcache.ErrCacheMiss) && firstErr == nil {
firstErr = err
}
}
for _, key := range keys {
_ = m.client.Delete(key)
// Also delete the tags key for this cache key
tagsKey := fmt.Sprintf("cache:tags:%s", key)
_ = m.client.Delete(tagsKey)
note(m.client.Delete(memcacheKey(key)))
note(m.client.Delete(memcacheTagKey("cache:tags:", key)))
}
// Delete the tag key itself
_ = m.client.Delete(tagKey)
return nil
if firstErr != nil {
return firstErr // keep the tag index so the invalidation can be retried
}
note(m.client.Delete(tagKey))
return firstErr
}
// DeleteByPattern removes all keys matching the pattern.
// Note: Memcache does not support pattern-based deletion natively.
// This is a no-op for memcache and returns an error.
// DeleteByPattern is not supported by Memcache; it always returns an error.
// Use tags (SetWithTags / DeleteByTag) for group invalidation instead.
func (m *MemcacheProvider) DeleteByPattern(ctx context.Context, pattern string) error {
return fmt.Errorf("pattern-based deletion is not supported by Memcache")
}
// Clear removes all items from the cache.
// It runs flush_all on every configured server and therefore requires MemcacheConfig.AllowFlush.
func (m *MemcacheProvider) Clear(ctx context.Context) error {
if !m.allowFlush {
return ErrFlushNotAllowed
}
return m.client.FlushAll()
}
// Exists checks if a key exists in the cache.
func (m *MemcacheProvider) Exists(ctx context.Context, key string) bool {
_, err := m.client.Get(key)
if ctx.Err() != nil {
return false
}
_, err := m.client.Get(memcacheKey(key))
return err == nil
}
// Close closes the provider and releases any resources.
// Close closes the provider and releases idle connections.
func (m *MemcacheProvider) Close() error {
// Memcache client doesn't have a close method
return nil
return m.client.Close()
}
// Stats returns statistics about the cache provider.
+114 -105
View File
@@ -2,20 +2,39 @@ package cache
import (
"context"
"errors"
"fmt"
"regexp"
"sync"
"sync/atomic"
"time"
"github.com/bitechdev/ResolveSpec/pkg/logger"
)
const (
defaultMemoryMaxSize = 10000
defaultMemoryCleanupInterval = time.Minute
)
// ErrClosed is returned by a provider that has been closed.
var ErrClosed = errors.New("cache: provider closed")
// memoryItem represents a cached item in memory.
type memoryItem struct {
Value []byte
Expiration time.Time
LastAccess time.Time
HitCount int64
Tags []string
lastAccess atomic.Int64 // unix nanos
hitCount atomic.Int64
}
func newMemoryItem(value []byte, expiration time.Time, tags []string) *memoryItem {
buf := make([]byte, len(value))
copy(buf, value)
item := &memoryItem{Value: buf, Expiration: expiration, Tags: tags}
item.lastAccess.Store(time.Now().UnixNano())
return item
}
// isExpired checks if the item has expired.
@@ -34,30 +53,72 @@ type MemoryProvider struct {
options *Options
hits atomic.Int64
misses atomic.Int64
closed bool
done chan struct{}
closeOnce sync.Once
}
// NewMemoryProvider creates a new in-memory cache provider.
// A MaxSize <= 0 selects the default (10000); use MaxSize -1 for an unbounded cache.
// A background goroutine removes expired items until Close is called.
func NewMemoryProvider(opts *Options) *MemoryProvider {
if opts == nil {
opts = &Options{
DefaultTTL: 5 * time.Minute,
MaxSize: 10000,
}
var o Options
if opts != nil {
o = *opts // do not mutate the caller's struct
} else {
o = Options{DefaultTTL: 5 * time.Minute}
}
if o.MaxSize == 0 {
o.MaxSize = defaultMemoryMaxSize
}
if o.CleanupInterval <= 0 {
o.CleanupInterval = defaultMemoryCleanupInterval
}
return &MemoryProvider{
m := &MemoryProvider{
items: make(map[string]*memoryItem),
tagToKeys: make(map[string]map[string]struct{}),
options: opts,
options: &o,
done: make(chan struct{}),
}
go m.janitor(o.CleanupInterval)
return m
}
func (m *MemoryProvider) janitor(interval time.Duration) {
defer logger.CatchPanic("cache.MemoryProvider.janitor")()
t := time.NewTicker(interval)
defer t.Stop()
for {
select {
case <-m.done:
return
case <-t.C:
m.CleanExpired(context.Background())
}
}
}
// removeLocked deletes a key and its tag associations. Caller must hold m.mu for writing.
func (m *MemoryProvider) removeLocked(key string) {
if item, ok := m.items[key]; ok {
for _, tag := range item.Tags {
if ks := m.tagToKeys[tag]; ks != nil {
delete(ks, key)
if len(ks) == 0 {
delete(m.tagToKeys, tag)
}
}
}
}
delete(m.items, key)
}
// Get retrieves a value from the cache by key.
func (m *MemoryProvider) Get(ctx context.Context, key string) ([]byte, bool) {
// First try with read lock for fast path
m.mu.RLock()
item, exists := m.items[key]
if !exists {
if !exists || m.closed {
m.mu.RUnlock()
m.misses.Add(1)
return nil, false
@@ -65,56 +126,29 @@ func (m *MemoryProvider) Get(ctx context.Context, key string) ([]byte, bool) {
if item.isExpired() {
m.mu.RUnlock()
// Upgrade to write lock to delete expired item
// Delete only if the entry is still the same expired one
m.mu.Lock()
delete(m.items, key)
if cur, ok := m.items[key]; ok && cur == item {
m.removeLocked(key)
}
m.mu.Unlock()
m.misses.Add(1)
return nil, false
}
// Update stats and access time with write lock
value := item.Value
item.lastAccess.Store(time.Now().UnixNano())
item.hitCount.Add(1)
out := make([]byte, len(item.Value))
copy(out, item.Value)
m.mu.RUnlock()
// Update access tracking with write lock
m.mu.Lock()
item.LastAccess = time.Now()
item.HitCount++
m.mu.Unlock()
m.hits.Add(1)
return value, true
return out, true
}
// Set stores a value in the cache with the specified TTL.
func (m *MemoryProvider) Set(ctx context.Context, key string, value []byte, ttl time.Duration) error {
m.mu.Lock()
defer m.mu.Unlock()
if ttl == 0 {
ttl = m.options.DefaultTTL
}
var expiration time.Time
if ttl > 0 {
expiration = time.Now().Add(ttl)
}
// Check max size and evict if necessary
if m.options.MaxSize > 0 && len(m.items) >= m.options.MaxSize {
if _, exists := m.items[key]; !exists {
m.evictOne()
}
}
m.items[key] = &memoryItem{
Value: value,
Expiration: expiration,
LastAccess: time.Now(),
}
return nil
return m.SetWithTags(ctx, key, value, ttl, nil)
}
// SetWithTags stores a value in the cache with the specified TTL and tags.
@@ -122,6 +156,10 @@ func (m *MemoryProvider) SetWithTags(ctx context.Context, key string, value []by
m.mu.Lock()
defer m.mu.Unlock()
if m.closed {
return ErrClosed
}
if ttl == 0 {
ttl = m.options.DefaultTTL
}
@@ -131,34 +169,14 @@ func (m *MemoryProvider) SetWithTags(ctx context.Context, key string, value []by
expiration = time.Now().Add(ttl)
}
// Check max size and evict if necessary
if m.options.MaxSize > 0 && len(m.items) >= m.options.MaxSize {
if _, exists := m.items[key]; !exists {
m.evictOne()
}
if _, exists := m.items[key]; exists {
m.removeLocked(key) // drops old tag associations
} else if m.options.MaxSize > 0 && len(m.items) >= m.options.MaxSize {
m.evictOne()
}
// Remove old tag associations if key exists
if oldItem, exists := m.items[key]; exists {
for _, tag := range oldItem.Tags {
if keySet, ok := m.tagToKeys[tag]; ok {
delete(keySet, key)
if len(keySet) == 0 {
delete(m.tagToKeys, tag)
}
}
}
}
m.items[key] = newMemoryItem(value, expiration, tags)
// Store the item
m.items[key] = &memoryItem{
Value: value,
Expiration: expiration,
LastAccess: time.Now(),
Tags: tags,
}
// Add new tag associations
for _, tag := range tags {
if m.tagToKeys[tag] == nil {
m.tagToKeys[tag] = make(map[string]struct{})
@@ -174,19 +192,7 @@ func (m *MemoryProvider) Delete(ctx context.Context, key string) error {
m.mu.Lock()
defer m.mu.Unlock()
// Remove tag associations
if item, exists := m.items[key]; exists {
for _, tag := range item.Tags {
if keySet, ok := m.tagToKeys[tag]; ok {
delete(keySet, key)
if len(keySet) == 0 {
delete(m.tagToKeys, tag)
}
}
}
}
delete(m.items, key)
m.removeLocked(key)
return nil
}
@@ -195,16 +201,13 @@ func (m *MemoryProvider) DeleteByTag(ctx context.Context, tag string) error {
m.mu.Lock()
defer m.mu.Unlock()
// Get all keys associated with this tag
keySet, exists := m.tagToKeys[tag]
if !exists {
return nil // No keys with this tag
}
// Delete all items with this tag
for key := range keySet {
if item, ok := m.items[key]; ok {
// Remove this tag from the item's tag list
newTags := make([]string, 0, len(item.Tags))
for _, t := range item.Tags {
if t != tag {
@@ -212,8 +215,7 @@ func (m *MemoryProvider) DeleteByTag(ctx context.Context, tag string) error {
}
}
// If item has no more tags, delete it
// Otherwise update its tags
// If item has no more tags, delete it; otherwise update its tags
if len(newTags) == 0 {
delete(m.items, key)
} else {
@@ -222,24 +224,24 @@ func (m *MemoryProvider) DeleteByTag(ctx context.Context, tag string) error {
}
}
// Remove the tag mapping
delete(m.tagToKeys, tag)
return nil
}
// DeleteByPattern removes all keys matching the pattern.
// The pattern is a Go regular expression (unanchored); it is compiled before the lock is taken.
func (m *MemoryProvider) DeleteByPattern(ctx context.Context, pattern string) error {
m.mu.Lock()
defer m.mu.Unlock()
re, err := regexp.Compile(pattern)
if err != nil {
return fmt.Errorf("invalid pattern: %w", err)
}
m.mu.Lock()
defer m.mu.Unlock()
for key := range m.items {
if re.MatchString(key) {
delete(m.items, key)
m.removeLocked(key)
}
}
@@ -252,6 +254,7 @@ func (m *MemoryProvider) Clear(ctx context.Context) error {
defer m.mu.Unlock()
m.items = make(map[string]*memoryItem)
m.tagToKeys = make(map[string]map[string]struct{})
m.hits.Store(0)
m.misses.Store(0)
return nil
@@ -270,12 +273,17 @@ func (m *MemoryProvider) Exists(ctx context.Context, key string) bool {
return !item.isExpired()
}
// Close closes the provider and releases any resources.
// Close closes the provider, stops the janitor and releases stored items.
// Later writes return ErrClosed and reads report a miss.
func (m *MemoryProvider) Close() error {
m.closeOnce.Do(func() { close(m.done) })
m.mu.Lock()
defer m.mu.Unlock()
m.items = nil
m.closed = true
m.items = make(map[string]*memoryItem)
m.tagToKeys = make(map[string]map[string]struct{})
return nil
}
@@ -284,7 +292,7 @@ func (m *MemoryProvider) Stats(ctx context.Context) (*CacheStats, error) {
m.mu.RLock()
defer m.mu.RUnlock()
// Clean expired items first
// Count non-expired items (read-only)
validKeys := 0
for _, item := range m.items {
if !item.isExpired() {
@@ -304,24 +312,25 @@ func (m *MemoryProvider) Stats(ctx context.Context) (*CacheStats, error) {
}
// evictOne removes one item from the cache using LRU strategy.
// Note: this is an O(n) scan. Caller must hold m.mu for writing.
func (m *MemoryProvider) evictOne() {
var oldestKey string
var oldestTime time.Time
var oldest int64
for key, item := range m.items {
if item.isExpired() {
delete(m.items, key)
m.removeLocked(key)
return
}
if oldestKey == "" || item.LastAccess.Before(oldestTime) {
if la := item.lastAccess.Load(); oldestKey == "" || la < oldest {
oldestKey = key
oldestTime = item.LastAccess
oldest = la
}
}
if oldestKey != "" {
delete(m.items, oldestKey)
m.removeLocked(oldestKey)
}
}
@@ -333,7 +342,7 @@ func (m *MemoryProvider) CleanExpired(ctx context.Context) int {
count := 0
for key, item := range m.items {
if item.isExpired() {
delete(m.items, key)
m.removeLocked(key)
count++
}
}
+51 -15
View File
@@ -3,15 +3,19 @@ package cache
import (
"context"
"fmt"
"strconv"
"strings"
"time"
"github.com/bitechdev/ResolveSpec/pkg/logger"
"github.com/redis/go-redis/v9"
)
// RedisProvider is a Redis implementation of the Provider interface.
type RedisProvider struct {
client *redis.Client
options *Options
client *redis.Client
options *Options
allowFlush bool
}
// RedisConfig contains Redis-specific configuration.
@@ -33,16 +37,25 @@ type RedisConfig struct {
// Options contains general cache options
Options *Options
// AllowFlush permits Clear() to run FLUSHDB, which wipes the entire logical Redis DB
// (including data that is not owned by this cache). Off by default.
AllowFlush bool
}
// NewRedisProvider creates a new Redis cache provider.
func NewRedisProvider(config *RedisConfig) (*RedisProvider, error) {
if config == nil {
config = &RedisConfig{
Host: "localhost",
Port: 6379,
DB: 0,
}
// Work on a copy so the caller's struct is not mutated
var cfg RedisConfig
if config != nil {
cfg = *config
} else {
cfg = RedisConfig{Host: "localhost", Port: 6379, DB: 0}
}
config = &cfg
if config.Options != nil {
o := *config.Options
config.Options = &o
}
if config.Host == "" {
@@ -77,8 +90,9 @@ func NewRedisProvider(config *RedisConfig) (*RedisProvider, error) {
}
return &RedisProvider{
client: client,
options: config.Options,
client: client,
options: config.Options,
allowFlush: config.AllowFlush,
}, nil
}
@@ -89,6 +103,8 @@ func (r *RedisProvider) Get(ctx context.Context, key string) ([]byte, bool) {
return nil, false
}
if err != nil {
// Reported as a miss (the Provider interface cannot express errors), but not silently
logger.Warn("cache: redis GET failed: %v", err)
return nil, false
}
return val, true
@@ -194,7 +210,7 @@ func (r *RedisProvider) DeleteByTag(ctx context.Context, tag string) error {
// DeleteByPattern removes all keys matching the pattern.
func (r *RedisProvider) DeleteByPattern(ctx context.Context, pattern string) error {
iter := r.client.Scan(ctx, 0, pattern, 0).Iterator()
iter := r.client.Scan(ctx, 0, pattern, 500).Iterator()
pipe := r.client.Pipeline()
count := 0
@@ -225,7 +241,11 @@ func (r *RedisProvider) DeleteByPattern(ctx context.Context, pattern string) err
}
// Clear removes all items from the cache.
// It runs FLUSHDB and therefore requires RedisConfig.AllowFlush.
func (r *RedisProvider) Clear(ctx context.Context) error {
if !r.allowFlush {
return ErrFlushNotAllowed
}
return r.client.FlushDB(ctx).Err()
}
@@ -244,8 +264,9 @@ func (r *RedisProvider) Close() error {
}
// Stats returns statistics about the cache provider.
// Only an allowlist of numeric counters from INFO is exposed, not the raw output.
func (r *RedisProvider) Stats(ctx context.Context) (*CacheStats, error) {
info, err := r.client.Info(ctx, "stats", "keyspace").Result()
info, err := r.client.Info(ctx, "stats").Result()
if err != nil {
return nil, fmt.Errorf("failed to get Redis stats: %w", err)
}
@@ -255,13 +276,28 @@ func (r *RedisProvider) Stats(ctx context.Context) (*CacheStats, error) {
return nil, fmt.Errorf("failed to get DB size: %w", err)
}
// Parse stats from INFO command
// This is a simplified version - you may want to parse more detailed stats
counters := map[string]int64{}
for _, line := range strings.Split(info, "\n") {
k, v, ok := strings.Cut(strings.TrimSpace(line), ":")
if !ok {
continue
}
switch k {
case "keyspace_hits", "keyspace_misses", "evicted_keys", "expired_keys":
if n, err := strconv.ParseInt(v, 10, 64); err == nil {
counters[k] = n
}
}
}
stats := &CacheStats{
Hits: counters["keyspace_hits"],
Misses: counters["keyspace_misses"],
Keys: dbSize,
ProviderType: "redis",
ProviderStats: map[string]any{
"info": info,
"evicted_keys": counters["evicted_keys"],
"expired_keys": counters["expired_keys"],
},
}