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
+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.