From 5af2ebd5bd6398ba96a269356fb05e706496d0a2 Mon Sep 17 00:00:00 2001 From: Mathieu Fenniak Date: Sat, 15 Nov 2025 18:58:47 -0700 Subject: [PATCH] fix: prevent Remove(key)...Get*(key) from returning a value computed before the Remove(key) --- modules/cache/cache.go | 169 +++++++++++++++++++++++------------------ 1 file changed, 93 insertions(+), 76 deletions(-) diff --git a/modules/cache/cache.go b/modules/cache/cache.go index 9bf4e9a00e..cdd9179623 100644 --- a/modules/cache/cache.go +++ b/modules/cache/cache.go @@ -9,6 +9,7 @@ import ( "strconv" "time" + "forgejo.org/modules/log" "forgejo.org/modules/setting" mc "code.forgejo.org/go-chi/cache" @@ -16,7 +17,11 @@ import ( _ "code.forgejo.org/go-chi/cache/memcache" // memcache plugin for cache ) -var conn mc.Cache +var ( + conn mc.Cache + ErrInconvertible = errors.New("value from cache was not convertible to expected type") + mutexMap MutexMap +) func newCache(cacheConfig setting.Cache) (mc.Cache, error) { return mc.NewCacher(mc.Options{ @@ -78,102 +83,104 @@ func GetCache() mc.Cache { return conn } -// GetString returns the key value from cache with callback when no key exists in cache -func GetString(key string, getFunc func() (string, error)) (string, error) { +// concurrencySafeGet is a single-process concurrency safe fetch from the cache, which provides the guarantee that after +// calling `cache.Remove(key)` and then `cache.Get*(key, ...)`, the value returned from cache will never have been +// computed **before** the `Remove` invocation. It uses in-memory synchronization, so its guarantee does not extend to +// a clustered configuration. +// +// getFunc is the computation for the value if caching is not available. convertFunc converts the cached value into the +// target type, and can return `ErrInconvertible` to indicate that the value couldn't be converts and should be +// recomputed instead; other errors are passed through. +func concurrencySafeGet[T any](key string, getFunc func() (T, error), convertFunc func(v any) (T, error)) (T, error) { if conn == nil || setting.CacheService.TTL <= 0 { return getFunc() } + // Use a double-checking method -- once before acquiring the write lock on this key (this block), and then again + // afterwards to avoid calling `getFunc` if it was computed while we were acquiring the lock. This causes two cache + // hits as a trade-off to minimize the number of lock acquisitions. If this trade-off causes too much cache load, + // this first `Get` could be removed -- the second one is performance-critical to ensure that after waiting a "long + // time" to compute w/ `getFunc`, we don't immediately redo that work after acquiring the lock. cached := conn.Get(key) - - if cached == nil { - value, err := getFunc() - if err != nil { - return value, err + if cached != nil { + retval, err := convertFunc(cached) + if err == nil { + return retval, nil + } else if !errors.Is(err, ErrInconvertible) { // for ErrInconvertible we'll fall through to recalculating the value + var zero T + return zero, err } - return value, conn.Put(key, value, setting.CacheService.TTLSeconds()) } - if value, ok := cached.(string); ok { - return value, nil + defer mutexMap.Lock(key)() + + // The second, performance-critical, check if the cache contains the target value. + cached = conn.Get(key) + if cached != nil { + retval, err := convertFunc(cached) + if err == nil { + return retval, nil + } else if !errors.Is(err, ErrInconvertible) { // for ErrInconvertible we'll fall through to recalculating the value + var zero T + return zero, err + } } - if stringer, ok := cached.(fmt.Stringer); ok { - return stringer.String(), nil + value, err := getFunc() + if err != nil { + return value, err } + return value, conn.Put(key, value, setting.CacheService.TTLSeconds()) +} - return fmt.Sprintf("%s", cached), nil +// GetString returns the key value from cache with callback when no key exists in cache +func GetString(key string, getFunc func() (string, error)) (string, error) { + v, err := concurrencySafeGet(key, getFunc, func(cached any) (string, error) { + if value, ok := cached.(string); ok { + return value, nil + } + if stringer, ok := cached.(fmt.Stringer); ok { + return stringer.String(), nil + } + return fmt.Sprintf("%s", cached), nil + }) + return v, err } // GetInt returns key value from cache with callback when no key exists in cache func GetInt(key string, getFunc func() (int, error)) (int, error) { - if conn == nil || setting.CacheService.TTL <= 0 { - return getFunc() - } - - cached := conn.Get(key) - - if cached == nil { - value, err := getFunc() - if err != nil { - return value, err + v, err := concurrencySafeGet(key, getFunc, func(cached any) (int, error) { + switch v := cached.(type) { + case int: + return v, nil + case string: + value, err := strconv.Atoi(v) + if err != nil { + return 0, err + } + return value, nil } - - return value, conn.Put(key, value, setting.CacheService.TTLSeconds()) - } - - switch v := cached.(type) { - case int: - return v, nil - case string: - value, err := strconv.Atoi(v) - if err != nil { - return 0, err - } - return value, nil - default: - value, err := getFunc() - if err != nil { - return value, err - } - return value, conn.Put(key, value, setting.CacheService.TTLSeconds()) - } + return 0, ErrInconvertible + }) + return v, err } // GetInt64 returns key value from cache with callback when no key exists in cache func GetInt64(key string, getFunc func() (int64, error)) (int64, error) { - if conn == nil || setting.CacheService.TTL <= 0 { - return getFunc() - } - - cached := conn.Get(key) - - if cached == nil { - value, err := getFunc() - if err != nil { - return value, err + v, err := concurrencySafeGet(key, getFunc, func(cached any) (int64, error) { + switch v := cached.(type) { + case int64: + return v, nil + case string: + value, err := strconv.ParseInt(v, 10, 64) + if err != nil { + return 0, err + } + return value, nil } - - return value, conn.Put(key, value, setting.CacheService.TTLSeconds()) - } - - switch v := conn.Get(key).(type) { - case int64: - return v, nil - case string: - value, err := strconv.ParseInt(v, 10, 64) - if err != nil { - return 0, err - } - return value, nil - default: - value, err := getFunc() - if err != nil { - return value, err - } - - return value, conn.Put(key, value, setting.CacheService.TTLSeconds()) - } + return 0, ErrInconvertible + }) + return v, err } // Remove key from cache @@ -181,5 +188,15 @@ func Remove(key string) { if conn == nil { return } - _ = conn.Delete(key) + + // The goal of `Remove(key)` is to ensure that *after* it is completed, a new value is computed. It's possible that + // a value is being computed for the key *right now* -- `getFunc` is about to return, we're about to delete the key, + // and then it will be Put into the cache with an out-of-date value computed before the `Remove(key)`. To prevent + // this we need the `Remove(key)` to also lock on the key, just like `Get*(key, ...)` does when computing it. + defer mutexMap.Lock(key)() + + err := conn.Delete(key) + if err != nil { + log.Error("unexpected error deleting key %s from cache: %v", err) + } }