mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-09-30 14:26:36 +08:00
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:
@@ -6,26 +6,30 @@ package config
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||
)
|
||||
|
||||
func TestListVisibleSystemConfigsUsesRedisCache(t *testing.T) {
|
||||
func TestListVisibleSystemConfigsUsesStoreCache(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
|
||||
repository.ResetSystemConfigRAMCacheForTest()
|
||||
if err := repository.InvalidateVisibleSystemConfigsCache(ctx); err != nil {
|
||||
t.Fatalf("InvalidateVisibleSystemConfigsCache() error = %v", err)
|
||||
}
|
||||
|
||||
// Warm cache
|
||||
if _, err := repository.ListVisibleSystemConfigs(ctx); err != nil {
|
||||
t.Fatalf("ListVisibleSystemConfigs() warm cache error = %v", err)
|
||||
}
|
||||
|
||||
// Directly insert a new config in DB (bypassing caching layer)
|
||||
if err := dbConn.Create(&model.SystemConfig{
|
||||
Key: "cache_probe_public_key",
|
||||
Value: "cache_probe_public_value",
|
||||
@@ -36,6 +40,7 @@ func TestListVisibleSystemConfigsUsesRedisCache(t *testing.T) {
|
||||
t.Fatalf("Create(cache_probe_public_key) error = %v", err)
|
||||
}
|
||||
|
||||
// Cached call: shouldn't return the new key yet
|
||||
cached, err := repository.ListVisibleSystemConfigs(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("ListVisibleSystemConfigs() cached call error = %v", err)
|
||||
@@ -46,18 +51,15 @@ func TestListVisibleSystemConfigsUsesRedisCache(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
exists, err := db.Redis.Exists(ctx, db.PrefixedKey(repository.SystemConfigVisibleListRedisKey)).Result()
|
||||
if err != nil {
|
||||
t.Fatalf("Redis.Exists() error = %v", err)
|
||||
}
|
||||
if exists == 0 {
|
||||
t.Fatal("expected visible config list cache key to exist")
|
||||
}
|
||||
|
||||
// Invalidate: triggers reload
|
||||
if err := repository.InvalidateVisibleSystemConfigsCache(ctx); err != nil {
|
||||
t.Fatalf("InvalidateVisibleSystemConfigsCache() error = %v", err)
|
||||
}
|
||||
|
||||
// Wait for Pub/Sub delivery in test environment
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
|
||||
// Refreshed call: should return the new key
|
||||
refreshed, err := repository.ListVisibleSystemConfigs(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("ListVisibleSystemConfigs() refreshed call error = %v", err)
|
||||
|
||||
@@ -5,10 +5,8 @@ package config
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"github.com/redis/go-redis/v9"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
@@ -21,6 +19,7 @@ func TestSystemConfigRAMCacheServesUntilInvalidated(t *testing.T) {
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
|
||||
repository.ResetSystemConfigRAMCacheForTest()
|
||||
if err := repository.InvalidateAllSystemConfigCaches(ctx); err != nil {
|
||||
t.Fatalf("InvalidateAllSystemConfigCaches() error = %v", err)
|
||||
@@ -34,15 +33,14 @@ func TestSystemConfigRAMCacheServesUntilInvalidated(t *testing.T) {
|
||||
t.Fatalf("GetSystemConfigByKey(site_name).Value = %q, want %q", warm.Value, "Wavelet")
|
||||
}
|
||||
|
||||
// Update DB directly (bypassing caching layer)
|
||||
if err := dbConn.Model(&model.SystemConfig{}).
|
||||
Where("key = ?", model.ConfigKeySiteName).
|
||||
Update("value", "ram_probe_value").Error; err != nil {
|
||||
t.Fatalf("Update(site_name) error = %v", err)
|
||||
}
|
||||
if err := db.HDel(ctx, repository.SystemConfigRedisHashKey, model.ConfigKeySiteName); err != nil {
|
||||
t.Fatalf("HDel(site_name) error = %v", err)
|
||||
}
|
||||
|
||||
// Should still return "Wavelet" since it's cached in RAM cache
|
||||
cached, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySiteName)
|
||||
if err != nil {
|
||||
t.Fatalf("GetSystemConfigByKey(site_name) cached error = %v", err)
|
||||
@@ -51,10 +49,15 @@ func TestSystemConfigRAMCacheServesUntilInvalidated(t *testing.T) {
|
||||
t.Fatalf("GetSystemConfigByKey(site_name).Value = %q, want stale RAM value %q", cached.Value, "Wavelet")
|
||||
}
|
||||
|
||||
// Invalidate the cache (triggers refresh callback)
|
||||
if err := repository.InvalidateSystemConfigCache(ctx, model.ConfigKeySiteName); err != nil {
|
||||
t.Fatalf("InvalidateSystemConfigCache(site_name) error = %v", err)
|
||||
}
|
||||
|
||||
// Allow some time for broadcast listener in test environment
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
|
||||
// Should return the updated value now
|
||||
refreshed, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySiteName)
|
||||
if err != nil {
|
||||
t.Fatalf("GetSystemConfigByKey(site_name) refreshed error = %v", err)
|
||||
@@ -62,47 +65,48 @@ func TestSystemConfigRAMCacheServesUntilInvalidated(t *testing.T) {
|
||||
if refreshed.Value != "ram_probe_value" {
|
||||
t.Fatalf("GetSystemConfigByKey(site_name).Value = %q, want %q", refreshed.Value, "ram_probe_value")
|
||||
}
|
||||
|
||||
exists, err := db.Redis.HExists(ctx, db.PrefixedKey(repository.SystemConfigRedisHashKey), model.ConfigKeySiteName).Result()
|
||||
if err != nil {
|
||||
t.Fatalf("HExists(site_name) error = %v", err)
|
||||
}
|
||||
if !exists {
|
||||
t.Fatal("expected redis hash field to be repopulated after refresh")
|
||||
}
|
||||
}
|
||||
|
||||
func TestInvalidateSystemConfigCacheClearsRedisField(t *testing.T) {
|
||||
func TestInvalidateSystemConfigCacheBroadcast(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySiteName)
|
||||
|
||||
repository.ResetSystemConfigRAMCacheForTest()
|
||||
|
||||
// Initially seed in cache
|
||||
_, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySiteName)
|
||||
if err != nil {
|
||||
t.Fatalf("GetSystemConfigByKey(site_name) error = %v", err)
|
||||
}
|
||||
_ = sc
|
||||
|
||||
// Update DB directly
|
||||
if err := dbConn.Model(&model.SystemConfig{}).
|
||||
Where("key = ?", model.ConfigKeySiteName).
|
||||
Update("value", "broadcast_value").Error; err != nil {
|
||||
t.Fatalf("Update(site_name) error = %v", err)
|
||||
}
|
||||
|
||||
// Invalidate: publishes to Redis and refreshes locally/other nodes
|
||||
if err := repository.InvalidateSystemConfigCache(ctx, model.ConfigKeySiteName); err != nil {
|
||||
t.Fatalf("InvalidateSystemConfigCache(site_name) error = %v", err)
|
||||
}
|
||||
|
||||
_, err = db.Redis.HGet(ctx, db.PrefixedKey(repository.SystemConfigRedisHashKey), model.ConfigKeySiteName).Result()
|
||||
if !errors.Is(err, redis.Nil) {
|
||||
t.Fatalf("HGet(site_name) error = %v, want redis.Nil", err)
|
||||
}
|
||||
|
||||
if err := dbConn.Model(&model.SystemConfig{}).
|
||||
Where("key = ?", model.ConfigKeySiteName).
|
||||
Update("value", "after_invalidate").Error; err != nil {
|
||||
t.Fatalf("Update(site_name) error = %v", err)
|
||||
}
|
||||
// Wait for Redis Pub/Sub delivery in test
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
|
||||
refreshed, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySiteName)
|
||||
if err != nil {
|
||||
t.Fatalf("GetSystemConfigByKey(site_name) refreshed error = %v", err)
|
||||
}
|
||||
if refreshed.Value != "after_invalidate" {
|
||||
t.Fatalf("GetSystemConfigByKey(site_name).Value = %q, want %q", refreshed.Value, "after_invalidate")
|
||||
if refreshed.Value != "broadcast_value" {
|
||||
t.Fatalf("GetSystemConfigByKey(site_name).Value = %q, want %q", refreshed.Value, "broadcast_value")
|
||||
}
|
||||
|
||||
// Verify Redis pub/sub channel received message
|
||||
if db.Redis != nil {
|
||||
// Just a sanity check: we can publish a new manual update and verify subscription triggers
|
||||
// which was done implicitly above.
|
||||
}
|
||||
}
|
||||
+148
-69
@@ -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()
|
||||
}
|
||||
|
||||
@@ -0,0 +1,235 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package oauth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/alicebob/miniredis/v2"
|
||||
"github.com/redis/go-redis/v9"
|
||||
"github.com/redis/go-redis/v9/maintnotifications"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
func setupOauthCacheTest(t *testing.T) (*miniredis.Miniredis, func()) {
|
||||
t.Helper()
|
||||
|
||||
miniRedis, err := miniredis.Run()
|
||||
if err != nil {
|
||||
t.Fatalf("failed to start miniredis: %v", err)
|
||||
}
|
||||
|
||||
db.Redis = redis.NewClient(&redis.Options{
|
||||
Addr: miniRedis.Addr(),
|
||||
MaintNotificationsConfig: &maintnotifications.Config{
|
||||
Mode: maintnotifications.ModeDisabled,
|
||||
},
|
||||
})
|
||||
|
||||
ResetOauthRAMCacheForTest()
|
||||
|
||||
cleanup := func() {
|
||||
StopOauthCacheListener()
|
||||
ResetOauthRAMCacheForTest()
|
||||
db.Redis.Close()
|
||||
miniRedis.Close()
|
||||
db.Redis = nil
|
||||
}
|
||||
return miniRedis, cleanup
|
||||
}
|
||||
|
||||
func TestTokenCache_GetSetInvalidate(t *testing.T) {
|
||||
_, cleanup := setupOauthCacheTest(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
tokenHash := "test-token-hash"
|
||||
token := &model.AccessToken{
|
||||
ID: 123,
|
||||
UserID: 456,
|
||||
TokenHash: tokenHash,
|
||||
Name: "test-token",
|
||||
}
|
||||
|
||||
// 1. Get from empty cache -> miss
|
||||
_, err := GetCachedToken(ctx, tokenHash)
|
||||
if err == nil {
|
||||
t.Fatal("expected cache miss for un-cached token")
|
||||
}
|
||||
|
||||
// 2. Set to cache
|
||||
SetCachedToken(ctx, tokenHash, token)
|
||||
|
||||
// 3. Get from cache -> hit
|
||||
cached, err := GetCachedToken(ctx, tokenHash)
|
||||
if err != nil {
|
||||
t.Fatalf("GetCachedToken() failed: %v", err)
|
||||
}
|
||||
if cached.ID != token.ID || cached.UserID != token.UserID {
|
||||
t.Fatalf("expected cached token %+v, got %+v", token, cached)
|
||||
}
|
||||
|
||||
// 4. Invalidate cache
|
||||
InvalidateCachedToken(ctx, tokenHash)
|
||||
|
||||
// 5. Get from cache -> miss
|
||||
_, err = GetCachedToken(ctx, tokenHash)
|
||||
if err == nil {
|
||||
t.Fatal("expected cache miss after invalidation")
|
||||
}
|
||||
}
|
||||
|
||||
func TestUserCache_GetSetInvalidate(t *testing.T) {
|
||||
_, cleanup := setupOauthCacheTest(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
userID := uint64(789)
|
||||
user := &model.User{
|
||||
ID: userID,
|
||||
Username: "testuser",
|
||||
Email: "test@example.com",
|
||||
}
|
||||
|
||||
// 1. Get from empty cache -> miss
|
||||
_, err := GetCachedUser(ctx, userID)
|
||||
if err == nil {
|
||||
t.Fatal("expected cache miss for un-cached user")
|
||||
}
|
||||
|
||||
// 2. Set to cache
|
||||
SetCachedUser(ctx, userID, user)
|
||||
|
||||
// 3. Get from cache -> hit
|
||||
cached, err := GetCachedUser(ctx, userID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetCachedUser() failed: %v", err)
|
||||
}
|
||||
if cached.ID != user.ID || cached.Username != user.Username {
|
||||
t.Fatalf("expected cached user %+v, got %+v", user, cached)
|
||||
}
|
||||
|
||||
// 4. Invalidate cache
|
||||
InvalidateCachedUser(ctx, userID)
|
||||
|
||||
// 5. Get from cache -> miss
|
||||
_, err = GetCachedUser(ctx, userID)
|
||||
if err == nil {
|
||||
t.Fatal("expected cache miss after invalidation")
|
||||
}
|
||||
}
|
||||
|
||||
func TestOauthCache_PubSubInvalidation(t *testing.T) {
|
||||
_, cleanup := setupOauthCacheTest(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
tokenHash := "pubsub-token-hash"
|
||||
token := &model.AccessToken{
|
||||
ID: 111,
|
||||
UserID: 222,
|
||||
TokenHash: tokenHash,
|
||||
}
|
||||
|
||||
userID := uint64(333)
|
||||
user := &model.User{
|
||||
ID: userID,
|
||||
Username: "pubsubuser",
|
||||
}
|
||||
|
||||
// Set caches so they are stored in RAM
|
||||
SetCachedToken(ctx, tokenHash, token)
|
||||
SetCachedUser(ctx, userID, user)
|
||||
|
||||
// Verify they are cached
|
||||
if _, ok := tokenRAM.GetIfPresent(tokenHash); !ok {
|
||||
t.Fatal("expected token to be in RAM cache")
|
||||
}
|
||||
if _, ok := userRAM.GetIfPresent(userID); !ok {
|
||||
t.Fatal("expected user to be in RAM cache")
|
||||
}
|
||||
|
||||
// Give Pub/Sub subscription time to establish
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
|
||||
// Publish invalidation messages directly to simulate peer node invalidation
|
||||
if err := db.Redis.Publish(ctx, oauthTokenInvalidationChannel, tokenHash).Err(); err != nil {
|
||||
t.Fatalf("publish token invalidation error: %v", err)
|
||||
}
|
||||
if err := db.Redis.Publish(ctx, oauthUserInvalidationChannel, "333").Err(); err != nil {
|
||||
t.Fatalf("publish user invalidation error: %v", err)
|
||||
}
|
||||
|
||||
// Wait for background pubsub handlers to process messages
|
||||
deadline := time.Now().Add(500 * time.Millisecond)
|
||||
for time.Now().Before(deadline) {
|
||||
_, tokenOk := tokenRAM.GetIfPresent(tokenHash)
|
||||
_, userOk := userRAM.GetIfPresent(userID)
|
||||
if !tokenOk && !userOk {
|
||||
break
|
||||
}
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
}
|
||||
|
||||
if _, ok := tokenRAM.GetIfPresent(tokenHash); ok {
|
||||
t.Fatal("expected token RAM cache to be invalidated by Pub/Sub")
|
||||
}
|
||||
if _, ok := userRAM.GetIfPresent(userID); ok {
|
||||
t.Fatal("expected user RAM cache to be invalidated by Pub/Sub")
|
||||
}
|
||||
}
|
||||
|
||||
func TestOauthCache_PubSubResetAll(t *testing.T) {
|
||||
_, cleanup := setupOauthCacheTest(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
tokenHash := "reset-token-hash"
|
||||
token := &model.AccessToken{
|
||||
ID: 444,
|
||||
UserID: 555,
|
||||
TokenHash: tokenHash,
|
||||
}
|
||||
|
||||
userID := uint64(666)
|
||||
user := &model.User{
|
||||
ID: userID,
|
||||
Username: "resetuser",
|
||||
}
|
||||
|
||||
SetCachedToken(ctx, tokenHash, token)
|
||||
SetCachedUser(ctx, userID, user)
|
||||
|
||||
// Give Pub/Sub subscription time to establish
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
|
||||
// Publish dynamic reset/wildcard to clear all
|
||||
if err := db.Redis.Publish(ctx, oauthTokenInvalidationChannel, "*").Err(); err != nil {
|
||||
t.Fatalf("publish token reset error: %v", err)
|
||||
}
|
||||
if err := db.Redis.Publish(ctx, oauthUserInvalidationChannel, "reset").Err(); err != nil {
|
||||
t.Fatalf("publish user reset error: %v", err)
|
||||
}
|
||||
|
||||
deadline := time.Now().Add(500 * time.Millisecond)
|
||||
for time.Now().Before(deadline) {
|
||||
_, tokenOk := tokenRAM.GetIfPresent(tokenHash)
|
||||
_, userOk := userRAM.GetIfPresent(userID)
|
||||
if !tokenOk && !userOk {
|
||||
break
|
||||
}
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
}
|
||||
|
||||
if _, ok := tokenRAM.GetIfPresent(tokenHash); ok {
|
||||
t.Fatal("expected token RAM cache to be fully cleared by '*'")
|
||||
}
|
||||
if _, ok := userRAM.GetIfPresent(userID); ok {
|
||||
t.Fatal("expected user RAM cache to be fully cleared by 'reset'")
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user