fix: harden concurrent quota and status updates
This commit is contained in:
+32
-65
@@ -274,27 +274,10 @@ func GetTokenById(id int) (*Token, error) {
|
||||
token := Token{Id: id}
|
||||
var err error = nil
|
||||
err = DB.First(&token, "id = ?", id).Error
|
||||
if shouldUpdateRedis(true, err) {
|
||||
gopool.Go(func() {
|
||||
if err := cacheSetToken(token); err != nil {
|
||||
common.SysLog("failed to update user status cache: " + err.Error())
|
||||
}
|
||||
})
|
||||
}
|
||||
return &token, err
|
||||
}
|
||||
|
||||
func GetTokenByKey(key string, fromDB bool) (token *Token, err error) {
|
||||
defer func() {
|
||||
// Update Redis cache asynchronously on successful DB read
|
||||
if shouldUpdateRedis(fromDB, err) && token != nil {
|
||||
gopool.Go(func() {
|
||||
if err := cacheSetToken(*token); err != nil {
|
||||
common.SysLog("failed to update user status cache: " + err.Error())
|
||||
}
|
||||
})
|
||||
}
|
||||
}()
|
||||
if !fromDB && common.RedisEnabled {
|
||||
// Try Redis first
|
||||
token, err := cacheGetTokenByKey(key)
|
||||
@@ -303,9 +286,18 @@ func GetTokenByKey(key string, fromDB bool) (token *Token, err error) {
|
||||
}
|
||||
// Don't return error - fall through to DB
|
||||
}
|
||||
fromDB = true
|
||||
err = DB.Where(commonKeyCol+" = ?", key).First(&token).Error
|
||||
return token, err
|
||||
token = &Token{}
|
||||
if err = DB.Where(commonKeyCol+" = ?", key).First(token).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if common.RedisEnabled {
|
||||
// 冷缓存时用数据库快照初始化;已存在的哈希只刷新 TTL,
|
||||
// 避免快照覆盖 Redis 中已被原子预扣的余额。初始化失败不影响本次读取。
|
||||
if _, cacheErr := cacheInitToken(*token); cacheErr != nil {
|
||||
common.SysLog("failed to init token cache: " + cacheErr.Error())
|
||||
}
|
||||
}
|
||||
return token, nil
|
||||
}
|
||||
|
||||
func (token *Token) Insert() error {
|
||||
@@ -316,47 +308,27 @@ func (token *Token) Insert() error {
|
||||
|
||||
// Update Make sure your token's fields is completed, because this will update non-zero values
|
||||
func (token *Token) Update() (err error) {
|
||||
err = DB.Model(token).Select("name", "status", "expired_time", "remain_quota", "unlimited_quota",
|
||||
"model_limits_enabled", "model_limits", "allow_ips", "group", "cross_group_retry", "auto_groups").Updates(token).Error
|
||||
if shouldUpdateRedis(true, err) {
|
||||
if cacheErr := cacheSetToken(*token); cacheErr != nil {
|
||||
common.SysLog("failed to update token cache: " + cacheErr.Error())
|
||||
if deleteErr := cacheDeleteToken(token.Key); deleteErr != nil {
|
||||
common.SysLog("failed to invalidate token cache after update: " + deleteErr.Error())
|
||||
}
|
||||
}
|
||||
// 写库前失效缓存并设置 fence,防止并发读者把过期快照重新写回缓存。
|
||||
if cacheErr := invalidateTokenCacheForMutation(token.Key); cacheErr != nil {
|
||||
common.SysLog("failed to invalidate token cache before update: " + cacheErr.Error())
|
||||
}
|
||||
return err
|
||||
return DB.Model(token).Select("name", "status", "expired_time", "remain_quota", "unlimited_quota",
|
||||
"model_limits_enabled", "model_limits", "allow_ips", "group", "cross_group_retry", "auto_groups").Updates(token).Error
|
||||
}
|
||||
|
||||
func (token *Token) SelectUpdate() (err error) {
|
||||
defer func() {
|
||||
if shouldUpdateRedis(true, err) {
|
||||
gopool.Go(func() {
|
||||
err := cacheSetToken(*token)
|
||||
if err != nil {
|
||||
common.SysLog("failed to update token cache: " + err.Error())
|
||||
}
|
||||
})
|
||||
}
|
||||
}()
|
||||
if cacheErr := invalidateTokenCacheForMutation(token.Key); cacheErr != nil {
|
||||
common.SysLog("failed to invalidate token cache before status update: " + cacheErr.Error())
|
||||
}
|
||||
// This can update zero values
|
||||
return DB.Model(token).Select("accessed_time", "status").Updates(token).Error
|
||||
}
|
||||
|
||||
func (token *Token) Delete() (err error) {
|
||||
defer func() {
|
||||
if shouldUpdateRedis(true, err) {
|
||||
gopool.Go(func() {
|
||||
err := cacheDeleteToken(token.Key)
|
||||
if err != nil {
|
||||
common.SysLog("failed to delete token cache: " + err.Error())
|
||||
}
|
||||
})
|
||||
}
|
||||
}()
|
||||
err = DB.Delete(token).Error
|
||||
return err
|
||||
if cacheErr := invalidateTokenCacheForMutation(token.Key); cacheErr != nil {
|
||||
common.SysLog("failed to invalidate token cache before delete: " + cacheErr.Error())
|
||||
}
|
||||
return DB.Delete(token).Error
|
||||
}
|
||||
|
||||
func (token *Token) IsModelLimitsEnabled() bool {
|
||||
@@ -408,8 +380,9 @@ func IncreaseTokenQuota(tokenId int, key string, quota int) (err error) {
|
||||
}
|
||||
if common.RedisEnabled {
|
||||
gopool.Go(func() {
|
||||
err := cacheIncrTokenQuota(key, int64(quota))
|
||||
if err != nil {
|
||||
// 守卫式增量:哈希不存在时跳过,由下次读取从数据库水合,
|
||||
// 绝不创建只有配额字段的残缺哈希。
|
||||
if _, err := cacheApplyTokenQuotaDelta(tokenId, key, int64(quota)); err != nil {
|
||||
common.SysLog("failed to increase token quota: " + err.Error())
|
||||
}
|
||||
})
|
||||
@@ -438,8 +411,7 @@ func DecreaseTokenQuota(id int, key string, quota int) (err error) {
|
||||
}
|
||||
if common.RedisEnabled {
|
||||
gopool.Go(func() {
|
||||
err := cacheDecrTokenQuota(key, int64(quota))
|
||||
if err != nil {
|
||||
if _, err := cacheApplyTokenQuotaDelta(id, key, int64(-quota)); err != nil {
|
||||
common.SysLog("failed to decrease token quota: " + err.Error())
|
||||
}
|
||||
})
|
||||
@@ -482,6 +454,9 @@ func BatchDeleteTokens(ids []int, userId int) (int, error) {
|
||||
tx.Rollback()
|
||||
return 0, err
|
||||
}
|
||||
if err := invalidateTokensCache(tokens); err != nil {
|
||||
common.SysLog("failed to invalidate token cache before batch delete: " + err.Error())
|
||||
}
|
||||
|
||||
if err := tx.Where("user_id = ? AND id IN (?)", userId, ids).Delete(&Token{}).Error; err != nil {
|
||||
tx.Rollback()
|
||||
@@ -492,14 +467,6 @@ func BatchDeleteTokens(ids []int, userId int) (int, error) {
|
||||
return 0, err
|
||||
}
|
||||
|
||||
if common.RedisEnabled {
|
||||
gopool.Go(func() {
|
||||
for _, t := range tokens {
|
||||
_ = cacheDeleteToken(t.Key)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
return len(tokens), nil
|
||||
}
|
||||
|
||||
@@ -540,7 +507,7 @@ func invalidateTokensCache(tokens []Token) error {
|
||||
if t.Key == "" {
|
||||
continue
|
||||
}
|
||||
if err := cacheDeleteToken(t.Key); err != nil && firstErr == nil {
|
||||
if err := invalidateTokenCacheForMutation(t.Key); err != nil && firstErr == nil {
|
||||
firstErr = err
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user