Files
CamTalk/backend/internal/store/cached_user.go

173 lines
5.2 KiB
Go
Raw Normal View History

package store
import (
"context"
"time"
"github.com/redis/go-redis/v9"
"github.com/hhs/camtalk/internal/trace"
)
// Redis key 前缀。
const (
refreshTokenPrefix = "auth:refresh:" // auth:refresh:{token_hash} → user_id
userRefreshPrefix = "auth:user_refresh:" // auth:user_refresh:{user_id} → Set of token_hash
)
// CachedUserRepository 装饰器,为 UserRepository 的 refresh token 操作增加 Redis 缓存。
// 读路径Redis miss → DB → 回填 Redis。
// 写路径:同步双写 Redis + DB。
// 删路径:同步双删 Redis + DB。
// Redis 操作失败时降级到纯 DB不阻断主流程。
type CachedUserRepository struct {
inner UserRepository
rdb *redis.Client
backfillTTL time.Duration // DB 回填 Redis 时使用的默认 TTL
}
// NewCachedUserRepository 创建带 Redis 缓存的 UserRepository 装饰器。
// backfillTTL: 从 DB 回填 Redis 时使用的 TTL因 DB 接口不返回 expiresAt
func NewCachedUserRepository(inner UserRepository, rdb *redis.Client, backfillTTL time.Duration) *CachedUserRepository {
if backfillTTL <= 0 {
backfillTTL = 24 * time.Hour
}
return &CachedUserRepository{
inner: inner,
rdb: rdb,
backfillTTL: backfillTTL,
}
}
// refreshTokenKey 生成 refresh token 的 Redis key。
func refreshTokenKey(tokenHash string) string {
return refreshTokenPrefix + tokenHash
}
// userRefreshKey 生成用户 refresh token 集合的 Redis key。
func userRefreshKey(userID string) string {
return userRefreshPrefix + userID
}
// --- 委托方法(不做缓存) ---
func (r *CachedUserRepository) Create(ctx context.Context, username, passwordHash string) (string, error) {
return r.inner.Create(ctx, username, passwordHash)
}
func (r *CachedUserRepository) FindByUsername(ctx context.Context, username string) (*User, error) {
return r.inner.FindByUsername(ctx, username)
}
func (r *CachedUserRepository) FindByID(ctx context.Context, id string) (*User, error) {
return r.inner.FindByID(ctx, id)
}
// --- 缓存方法 ---
// SaveRefreshToken Write-Through先写 DB再写 Redis。
func (r *CachedUserRepository) SaveRefreshToken(ctx context.Context, userID, tokenHash string, expiresAt time.Time) error {
// 先写 DB
if err := r.inner.SaveRefreshToken(ctx, userID, tokenHash, expiresAt); err != nil {
return err
}
// 写 RedisSET + SADD设置 TTL 为 token 剩余有效期
ttl := time.Until(expiresAt)
if ttl <= 0 {
return nil
}
key := refreshTokenKey(tokenHash)
pipe := r.rdb.Pipeline()
pipe.Set(ctx, key, userID, ttl)
pipe.SAdd(ctx, userRefreshKey(userID), tokenHash)
if _, err := pipe.Exec(ctx); err != nil {
log := trace.FromContext(ctx)
log.Warnw("redis cache write failed for refresh token", "error", err)
// 降级DB 已写入成功Redis 失败不影响正确性
}
return nil
}
// FindRefreshToken Read-Through先查 Redismiss 时查 DB 并回填。
func (r *CachedUserRepository) FindRefreshToken(ctx context.Context, tokenHash string) (string, error) {
key := refreshTokenKey(tokenHash)
// 查 Redis
userID, err := r.rdb.Get(ctx, key).Result()
if err == nil {
return userID, nil
}
// redis.Nil 表示 key 不存在,其他错误记录日志后降级到 DB
if err != redis.Nil {
log := trace.FromContext(ctx)
log.Warnw("redis cache read failed for refresh token", "error", err)
}
// 降级到 DB
userID, err = r.inner.FindRefreshToken(ctx, tokenHash)
if err != nil {
return "", err
}
// 回填 RedisSET + SADDTTL 使用保守默认值
go func() {
bgCtx := context.Background()
pipe := r.rdb.Pipeline()
pipe.Set(bgCtx, key, userID, r.backfillTTL)
pipe.SAdd(bgCtx, userRefreshKey(userID), tokenHash)
_, _ = pipe.Exec(bgCtx)
}()
return userID, nil
}
// DeleteRefreshToken 双删:先删 DB再删 Redis。
func (r *CachedUserRepository) DeleteRefreshToken(ctx context.Context, tokenHash string) error {
// 先从 Redis 获取 user_id用于从集合中移除
userID, _ := r.rdb.Get(ctx, refreshTokenKey(tokenHash)).Result()
// 删 DB
if err := r.inner.DeleteRefreshToken(ctx, tokenHash); err != nil {
return err
}
// 删 Redis
key := refreshTokenKey(tokenHash)
pipe := r.rdb.Pipeline()
pipe.Del(ctx, key)
if userID != "" {
pipe.SRem(ctx, userRefreshKey(userID), tokenHash)
}
if _, err := pipe.Exec(ctx); err != nil {
log := trace.FromContext(ctx)
log.Warnw("redis cache delete failed for refresh token", "error", err)
}
return nil
}
// DeleteUserRefreshTokens 批量清理:先从 Redis 获取集合,逐个删缓存,再删 DB。
func (r *CachedUserRepository) DeleteUserRefreshTokens(ctx context.Context, userID string) error {
userKey := userRefreshKey(userID)
// 从 Redis 获取该用户所有 token hash
hashes, _ := r.rdb.SMembers(ctx, userKey).Result()
// 批量删除 Redis 缓存
if len(hashes) > 0 {
keys := make([]string, 0, len(hashes)+1)
for _, h := range hashes {
keys = append(keys, refreshTokenKey(h))
}
keys = append(keys, userKey)
if err := r.rdb.Del(ctx, keys...).Err(); err != nil {
log := trace.FromContext(ctx)
log.Warnw("redis cache batch delete failed for user refresh tokens", "error", err, "user_id", userID)
}
}
// 删 DB无论 Redis 是否成功都执行)
return r.inner.DeleteUserRefreshTokens(ctx, userID)
}