fix: harden concurrent quota and status updates

This commit is contained in:
CaIon
2026-08-11 22:03:47 +08:00
parent 50e5377ea5
commit ccd535ef8e
14 changed files with 702 additions and 203 deletions
+32 -65
View File
@@ -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
}
}