diff --git a/audit/pkg/cache.audit.md b/audit/pkg/cache.audit.md index a6f984c..95aa558 100644 --- a/audit/pkg/cache.audit.md +++ b/audit/pkg/cache.audit.md @@ -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 diff --git a/pkg/cache/cache.go b/pkg/cache/cache.go index 71b6fc8..b5b2099 100644 --- a/pkg/cache/cache.go +++ b/pkg/cache/cache.go @@ -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 } diff --git a/pkg/cache/cache_manager.go b/pkg/cache/cache_manager.go index 2f88cda..4b8d919 100644 --- a/pkg/cache/cache_manager.go +++ b/pkg/cache/cache_manager.go @@ -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 diff --git a/pkg/cache/example_usage.go b/pkg/cache/example_usage.go index c6a545b..9dbc28e 100644 --- a/pkg/cache/example_usage.go +++ b/pkg/cache/example_usage.go @@ -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 ( diff --git a/pkg/cache/hardening_test.go b/pkg/cache/hardening_test.go new file mode 100644 index 0000000..e04750c --- /dev/null +++ b/pkg/cache/hardening_test.go @@ -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") + } +} diff --git a/pkg/cache/provider.go b/pkg/cache/provider.go index 87ba1e7..583f192 100644 --- a/pkg/cache/provider.go +++ b/pkg/cache/provider.go @@ -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 } diff --git a/pkg/cache/provider_memcache.go b/pkg/cache/provider_memcache.go index ab42837..da5af4a 100644 --- a/pkg/cache/provider_memcache.go +++ b/pkg/cache/provider_memcache.go @@ -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. diff --git a/pkg/cache/provider_memory.go b/pkg/cache/provider_memory.go index 1f70758..1fc5dba 100644 --- a/pkg/cache/provider_memory.go +++ b/pkg/cache/provider_memory.go @@ -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++ } } diff --git a/pkg/cache/provider_redis.go b/pkg/cache/provider_redis.go index f5c5ac4..fbd8236 100644 --- a/pkg/cache/provider_redis.go +++ b/pkg/cache/provider_redis.go @@ -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"], }, }