refactor(oauth): replace legacy oauth cache with standard ram cache and add pubsub synchronization

- Replaced custom map-based cache in apps/oauth/cache.go with standard pkg/cache/ram framework.
- Implemented Redis Pub/Sub invalidation channels for distributed token and user cache synchronization.
- Created apps/oauth/cache_test.go to verify local cache operations and pub/sub broadcasts.

refactor(cache): generic RAM cache with CoW and unified preheating

Replaced L2 Redis cache and old cache package with process-local generic pkg/cache/ram. Implemented Copy-on-Write for reads, fine-grained locks per type for writes, and unified preheating in bootstrap. Changed cache invalidation to lazy-loading to resolve SQLite deadlocks during transactions.
This commit is contained in:
ryan
2026-06-26 22:55:14 +08:00
parent a4f6c2ae34
commit 681de3b8cc
15 changed files with 972 additions and 251 deletions
+148 -69
View File
@@ -6,65 +6,35 @@ package oauth
import (
"context"
"fmt"
"strconv"
"sync"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/pkg/cache/ram"
)
type cacheEntry struct {
value any
expiredAt time.Time
}
type memoryCache struct {
sync.RWMutex
items map[string]cacheEntry
}
var localCache = &memoryCache{
items: make(map[string]cacheEntry),
}
func (c *memoryCache) Set(key string, val any, ttl time.Duration) {
c.Lock()
defer c.Unlock()
c.items[key] = cacheEntry{
value: val,
expiredAt: time.Now().Add(ttl),
}
}
func (c *memoryCache) Get(key string) (any, bool) {
c.RLock()
item, ok := c.items[key]
if !ok {
c.RUnlock()
return nil, false
}
if time.Now().After(item.expiredAt) {
c.RUnlock()
c.Lock()
if item, ok = c.items[key]; ok && time.Now().After(item.expiredAt) {
delete(c.items, key)
}
c.Unlock()
return nil, false
}
c.RUnlock()
return item.value, true
}
func (c *memoryCache) Delete(key string) {
c.Lock()
defer c.Unlock()
delete(c.items, key)
}
const (
tokenCacheTTL = 5 * time.Minute
userCacheTTL = 5 * time.Minute
//nolint:gosec // This is a Redis Pub/Sub channel name, not a credential
oauthTokenInvalidationChannel = "oauth:token_invalidation"
oauthUserInvalidationChannel = "oauth:user_invalidation"
)
var (
tokenRAM = ram.MustNew[string, *model.AccessToken](ram.Options{MaximumSize: 2048})
userRAM = ram.MustNew[uint64, *model.User](ram.Options{MaximumSize: 2048})
tokenListenerOnce sync.Once
tokenListenerCtx context.Context
tokenListenerCancel context.CancelFunc
userListenerOnce sync.Once
userListenerCtx context.Context
userListenerCancel context.CancelFunc
)
func tokenCacheKey(tokenHash string) string {
@@ -75,20 +45,98 @@ func userCacheKey(userID uint64) string {
return fmt.Sprintf("oauth:user:%d", userID)
}
func ensureTokenCacheListener() {
if db.Redis == nil {
return
}
tokenListenerOnce.Do(startTokenCacheInvalidationListener)
}
func startTokenCacheInvalidationListener() {
tokenListenerCtx, tokenListenerCancel = context.WithCancel(context.Background())
go func() {
pubsub := db.Redis.Subscribe(tokenListenerCtx, oauthTokenInvalidationChannel)
defer func() {
_ = pubsub.Close()
}()
go func() {
<-tokenListenerCtx.Done()
_ = pubsub.Close()
}()
for msg := range pubsub.Channel() {
tokenHash := msg.Payload
if tokenHash == "" || tokenHash == "*" || tokenHash == "reset" {
tokenRAM.InvalidateAll()
} else {
tokenRAM.Invalidate(tokenHash)
}
}
}()
}
func publishTokenRAMInvalidation(ctx context.Context, tokenHash string) {
if db.Redis == nil {
return
}
_ = db.Redis.Publish(ctx, oauthTokenInvalidationChannel, tokenHash).Err()
}
func ensureUserCacheListener() {
if db.Redis == nil {
return
}
userListenerOnce.Do(startUserCacheInvalidationListener)
}
func startUserCacheInvalidationListener() {
userListenerCtx, userListenerCancel = context.WithCancel(context.Background())
go func() {
pubsub := db.Redis.Subscribe(userListenerCtx, oauthUserInvalidationChannel)
defer func() {
_ = pubsub.Close()
}()
go func() {
<-userListenerCtx.Done()
_ = pubsub.Close()
}()
for msg := range pubsub.Channel() {
userIDStr := msg.Payload
if userIDStr == "" || userIDStr == "*" || userIDStr == "reset" {
userRAM.InvalidateAll()
} else if userID, err := strconv.ParseUint(userIDStr, 10, 64); err == nil {
userRAM.Invalidate(userID)
}
}
}()
}
func publishUserRAMInvalidation(ctx context.Context, userID uint64) {
if db.Redis == nil {
return
}
_ = db.Redis.Publish(ctx, oauthUserInvalidationChannel, strconv.FormatUint(userID, 10)).Err()
}
// GetCachedToken 获取缓存的 AccessToken
func GetCachedToken(ctx context.Context, tokenHash string) (*model.AccessToken, error) {
key := tokenCacheKey(tokenHash)
if val, ok := localCache.Get(key); ok {
if token, ok := val.(*model.AccessToken); ok {
return token, nil
}
ensureTokenCacheListener()
if val, ok := tokenRAM.GetIfPresent(tokenHash); ok {
return val, nil
}
if db.Redis != nil {
var token model.AccessToken
key := tokenCacheKey(tokenHash)
if err := db.GetJSON(ctx, key, &token); err == nil {
// Write back to local cache
localCache.Set(key, &token, tokenCacheTTL)
tokenRAM.Set(tokenHash, &token)
return &token, nil
}
}
@@ -97,36 +145,41 @@ func GetCachedToken(ctx context.Context, tokenHash string) (*model.AccessToken,
// SetCachedToken 设置 AccessToken 缓存
func SetCachedToken(ctx context.Context, tokenHash string, token *model.AccessToken) {
key := tokenCacheKey(tokenHash)
localCache.Set(key, token, tokenCacheTTL)
ensureTokenCacheListener()
tokenRAM.Set(tokenHash, token)
if db.Redis != nil {
key := tokenCacheKey(tokenHash)
_ = db.SetJSON(ctx, key, token, tokenCacheTTL)
}
}
// InvalidateCachedToken 吊销/删除 token 缓存
func InvalidateCachedToken(ctx context.Context, tokenHash string) {
key := tokenCacheKey(tokenHash)
localCache.Delete(key)
ensureTokenCacheListener()
tokenRAM.Invalidate(tokenHash)
if db.Redis != nil {
key := tokenCacheKey(tokenHash)
_ = db.Redis.Del(ctx, db.PrefixedKey(key)).Err()
publishTokenRAMInvalidation(ctx, tokenHash)
}
}
// GetCachedUser 获取缓存的 User
func GetCachedUser(ctx context.Context, userID uint64) (*model.User, error) {
key := userCacheKey(userID)
if val, ok := localCache.Get(key); ok {
if u, ok := val.(*model.User); ok {
return u, nil
}
ensureUserCacheListener()
if val, ok := userRAM.GetIfPresent(userID); ok {
return val, nil
}
if db.Redis != nil {
var u model.User
key := userCacheKey(userID)
if err := db.GetJSON(ctx, key, &u); err == nil {
// Write back to local cache
localCache.Set(key, &u, userCacheTTL)
userRAM.Set(userID, &u)
return &u, nil
}
}
@@ -135,18 +188,44 @@ func GetCachedUser(ctx context.Context, userID uint64) (*model.User, error) {
// SetCachedUser 设置 User 缓存
func SetCachedUser(ctx context.Context, userID uint64, u *model.User) {
key := userCacheKey(userID)
localCache.Set(key, u, userCacheTTL)
ensureUserCacheListener()
userRAM.Set(userID, u)
if db.Redis != nil {
key := userCacheKey(userID)
_ = db.SetJSON(ctx, key, u, userCacheTTL)
}
}
// InvalidateCachedUser 吊销/失效 User 缓存
func InvalidateCachedUser(ctx context.Context, userID uint64) {
key := userCacheKey(userID)
localCache.Delete(key)
ensureUserCacheListener()
userRAM.Invalidate(userID)
if db.Redis != nil {
key := userCacheKey(userID)
_ = db.Redis.Del(ctx, db.PrefixedKey(key)).Err()
publishUserRAMInvalidation(ctx, userID)
}
}
// StopOauthCacheListener stops both token and user Redis Pub/Sub subscription listeners and resets the sync.Once guards.
func StopOauthCacheListener() {
if tokenListenerCancel != nil {
tokenListenerCancel()
tokenListenerCancel = nil
}
tokenListenerOnce = sync.Once{}
if userListenerCancel != nil {
userListenerCancel()
userListenerCancel = nil
}
userListenerOnce = sync.Once{}
}
// ResetOauthRAMCacheForTest clears only the process-local RAM cache.
func ResetOauthRAMCacheForTest() {
tokenRAM.InvalidateAll()
userRAM.InvalidateAll()
}