Skip to content
Toggle navigation
P
Projects
G
Groups
S
Snippets
Help
phsl
/
new-api
This project
Loading...
Sign in
Toggle navigation
Go to a project
Project
Repository
Issues
0
Merge Requests
0
Pipelines
Wiki
Snippets
Members
Activity
Graph
Charts
Create a new issue
Jobs
Commits
Issue Boards
Files
Commits
Branches
Tags
Contributors
Graph
Compare
Charts
Unverified
Commit
ccd535ef
authored
Aug 10, 2026
by
CaIon
Browse files
Options
Browse Files
Download
Email Patches
Plain Diff
fix: harden concurrent quota and status updates
parent
50e5377e
Hide whitespace changes
Inline
Side-by-side
Showing
14 changed files
with
698 additions
and
199 deletions
+698
-199
controller/token.go
+2
-2
model/channel.go
+22
-11
model/channel_status_test.go
+102
-0
model/quota_reserve.go
+177
-5
model/quota_reserve_test.go
+228
-0
model/subscription.go
+13
-3
model/token.go
+32
-65
model/token_auto_groups_cache_test.go
+15
-4
model/token_cache.go
+75
-33
model/user.go
+2
-35
model/user_cache.go
+0
-30
service/billing_session.go
+11
-0
service/funding_source.go
+11
-1
service/quota.go
+8
-10
No files found.
controller/token.go
View file @
ccd535ef
...
@@ -279,7 +279,7 @@ func AddToken(c *gin.Context) {
...
@@ -279,7 +279,7 @@ func AddToken(c *gin.Context) {
common
.
ApiErrorI18n
(
c
,
i18n
.
MsgTokenQuotaNegative
)
common
.
ApiErrorI18n
(
c
,
i18n
.
MsgTokenQuotaNegative
)
return
return
}
}
maxQuotaValue
:=
int
((
1000000000
*
common
.
QuotaPerUnit
)
)
maxQuotaValue
:=
common
.
QuotaFromFloat
(
1000000000
*
common
.
QuotaPerUnit
)
if
token
.
RemainQuota
>
maxQuotaValue
{
if
token
.
RemainQuota
>
maxQuotaValue
{
common
.
ApiErrorI18n
(
c
,
i18n
.
MsgTokenQuotaExceedMax
,
map
[
string
]
any
{
"Max"
:
maxQuotaValue
})
common
.
ApiErrorI18n
(
c
,
i18n
.
MsgTokenQuotaExceedMax
,
map
[
string
]
any
{
"Max"
:
maxQuotaValue
})
return
return
...
@@ -373,7 +373,7 @@ func UpdateToken(c *gin.Context) {
...
@@ -373,7 +373,7 @@ func UpdateToken(c *gin.Context) {
common
.
ApiErrorI18n
(
c
,
i18n
.
MsgTokenQuotaNegative
)
common
.
ApiErrorI18n
(
c
,
i18n
.
MsgTokenQuotaNegative
)
return
return
}
}
maxQuotaValue
:=
int
((
1000000000
*
common
.
QuotaPerUnit
)
)
maxQuotaValue
:=
common
.
QuotaFromFloat
(
1000000000
*
common
.
QuotaPerUnit
)
if
token
.
RemainQuota
>
maxQuotaValue
{
if
token
.
RemainQuota
>
maxQuotaValue
{
common
.
ApiErrorI18n
(
c
,
i18n
.
MsgTokenQuotaExceedMax
,
map
[
string
]
any
{
"Max"
:
maxQuotaValue
})
common
.
ApiErrorI18n
(
c
,
i18n
.
MsgTokenQuotaExceedMax
,
map
[
string
]
any
{
"Max"
:
maxQuotaValue
})
return
return
...
...
model/channel.go
View file @
ccd535ef
...
@@ -346,11 +346,21 @@ func (channel *Channel) Save() error {
...
@@ -346,11 +346,21 @@ func (channel *Channel) Save() error {
return
DB
.
Save
(
channel
)
.
Error
return
DB
.
Save
(
channel
)
.
Error
}
}
func
(
channel
*
Channel
)
SaveWithoutKey
()
error
{
// saveStatusState persists only the fields owned by the channel status flow.
// Keeping this allowlist here prevents a stale channel snapshot from
// overwriting credentials, accounting counters, or channel configuration.
func
(
channel
*
Channel
)
saveStatusState
()
error
{
if
channel
.
Id
==
0
{
if
channel
.
Id
==
0
{
return
errors
.
New
(
"channel ID is 0"
)
return
errors
.
New
(
"channel ID is 0"
)
}
}
return
DB
.
Omit
(
"key"
)
.
Save
(
channel
)
.
Error
updates
:=
map
[
string
]
any
{
"status"
:
channel
.
Status
,
"other_info"
:
channel
.
OtherInfo
,
}
if
channel
.
ChannelInfo
.
IsMultiKey
{
updates
[
"channel_info"
]
=
channel
.
ChannelInfo
}
return
DB
.
Model
(
&
Channel
{})
.
Where
(
"id = ?"
,
channel
.
Id
)
.
Updates
(
updates
)
.
Error
}
}
func
GetAllChannels
(
startIdx
int
,
num
int
,
selectAll
bool
,
idSort
bool
,
sortOptions
...
ChannelSortOptions
)
([]
*
Channel
,
error
)
{
func
GetAllChannels
(
startIdx
int
,
num
int
,
selectAll
bool
,
idSort
bool
,
sortOptions
...
ChannelSortOptions
)
([]
*
Channel
,
error
)
{
...
@@ -713,19 +723,24 @@ func UpdateChannelStatus(channelId int, usingKey string, status int, reason stri
...
@@ -713,19 +723,24 @@ func UpdateChannelStatus(channelId int, usingKey string, status int, reason stri
if
common
.
MemoryCacheEnabled
{
if
common
.
MemoryCacheEnabled
{
channelStatusLock
.
Lock
()
channelStatusLock
.
Lock
()
defer
channelStatusLock
.
Unlock
()
defer
channelStatusLock
.
Unlock
()
}
// ChannelInfo stores both multi-key status and the polling cursor. Hold the
// same per-channel lock from the first read through persistence so neither
// writer can save a stale JSON snapshot over the other.
pollingLock
:=
GetChannelPollingLock
(
channelId
)
pollingLock
.
Lock
()
defer
pollingLock
.
Unlock
()
if
common
.
MemoryCacheEnabled
{
channelCache
,
_
:=
CacheGetChannel
(
channelId
)
channelCache
,
_
:=
CacheGetChannel
(
channelId
)
if
channelCache
==
nil
{
if
channelCache
==
nil
{
return
false
return
false
}
}
if
channelCache
.
ChannelInfo
.
IsMultiKey
{
if
channelCache
.
ChannelInfo
.
IsMultiKey
{
// Use per-channel lock to prevent concurrent map read/write with GetNextEnabledKey
beforeStatus
:=
channelCache
.
Status
beforeStatus
:=
channelCache
.
Status
pollingLock
:=
GetChannelPollingLock
(
channelId
)
pollingLock
.
Lock
()
// 如果是多Key模式,更新缓存中的状态
// 如果是多Key模式,更新缓存中的状态
handlerMultiKeyUpdate
(
channelCache
,
usingKey
,
status
,
reason
)
handlerMultiKeyUpdate
(
channelCache
,
usingKey
,
status
,
reason
)
pollingLock
.
Unlock
()
if
beforeStatus
!=
channelCache
.
Status
{
if
beforeStatus
!=
channelCache
.
Status
{
CacheUpdateChannelStatus
(
channelId
,
channelCache
.
Status
)
CacheUpdateChannelStatus
(
channelId
,
channelCache
.
Status
)
}
}
...
@@ -759,11 +774,7 @@ func UpdateChannelStatus(channelId int, usingKey string, status int, reason stri
...
@@ -759,11 +774,7 @@ func UpdateChannelStatus(channelId int, usingKey string, status int, reason stri
if
channel
.
ChannelInfo
.
IsMultiKey
{
if
channel
.
ChannelInfo
.
IsMultiKey
{
beforeStatus
:=
channel
.
Status
beforeStatus
:=
channel
.
Status
// Protect map writes with the same per-channel lock used by readers
pollingLock
:=
GetChannelPollingLock
(
channelId
)
pollingLock
.
Lock
()
handlerMultiKeyUpdate
(
channel
,
usingKey
,
status
,
reason
)
handlerMultiKeyUpdate
(
channel
,
usingKey
,
status
,
reason
)
pollingLock
.
Unlock
()
if
beforeStatus
!=
channel
.
Status
{
if
beforeStatus
!=
channel
.
Status
{
shouldUpdateAbilities
=
true
shouldUpdateAbilities
=
true
}
}
...
@@ -775,7 +786,7 @@ func UpdateChannelStatus(channelId int, usingKey string, status int, reason stri
...
@@ -775,7 +786,7 @@ func UpdateChannelStatus(channelId int, usingKey string, status int, reason stri
channel
.
Status
=
status
channel
.
Status
=
status
shouldUpdateAbilities
=
true
shouldUpdateAbilities
=
true
}
}
err
=
channel
.
SaveWithoutKey
()
err
=
channel
.
saveStatusState
()
if
err
!=
nil
{
if
err
!=
nil
{
common
.
SysLog
(
fmt
.
Sprintf
(
"failed to update channel status: channel_id=%d, status=%d, error=%v"
,
channel
.
Id
,
status
,
err
))
common
.
SysLog
(
fmt
.
Sprintf
(
"failed to update channel status: channel_id=%d, status=%d, error=%v"
,
channel
.
Id
,
status
,
err
))
return
false
return
false
...
...
model/channel_status_test.go
0 → 100644
View file @
ccd535ef
package
model
import
(
"testing"
"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/constant"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func
setupChannelStatusTest
(
t
*
testing
.
T
)
{
t
.
Helper
()
truncateTables
(
t
)
require
.
NoError
(
t
,
DB
.
Exec
(
"DELETE FROM abilities"
)
.
Error
)
require
.
NoError
(
t
,
DB
.
Exec
(
"DELETE FROM channels"
)
.
Error
)
memoryCacheEnabled
:=
common
.
MemoryCacheEnabled
common
.
MemoryCacheEnabled
=
false
t
.
Cleanup
(
func
()
{
common
.
MemoryCacheEnabled
=
memoryCacheEnabled
})
}
func
TestUpdateChannelStatusPersistsMultiKeyState
(
t
*
testing
.
T
)
{
setupChannelStatusTest
(
t
)
channel
:=
Channel
{
Name
:
"multi-key-status"
,
Key
:
"key-a
\n
key-b"
,
Status
:
common
.
ChannelStatusEnabled
,
ChannelInfo
:
ChannelInfo
{
IsMultiKey
:
true
,
MultiKeySize
:
2
,
MultiKeyMode
:
constant
.
MultiKeyModePolling
,
MultiKeyPollingIndex
:
1
,
},
}
require
.
NoError
(
t
,
DB
.
Create
(
&
channel
)
.
Error
)
changed
:=
UpdateChannelStatus
(
channel
.
Id
,
"key-a"
,
common
.
ChannelStatusAutoDisabled
,
"provider rejected key"
)
require
.
True
(
t
,
changed
)
var
stored
Channel
require
.
NoError
(
t
,
DB
.
First
(
&
stored
,
channel
.
Id
)
.
Error
)
assert
.
Equal
(
t
,
common
.
ChannelStatusEnabled
,
stored
.
Status
)
assert
.
Equal
(
t
,
common
.
ChannelStatusAutoDisabled
,
stored
.
ChannelInfo
.
MultiKeyStatusList
[
0
])
assert
.
Equal
(
t
,
"provider rejected key"
,
stored
.
ChannelInfo
.
MultiKeyDisabledReason
[
0
])
assert
.
NotZero
(
t
,
stored
.
ChannelInfo
.
MultiKeyDisabledTime
[
0
])
assert
.
Equal
(
t
,
1
,
stored
.
ChannelInfo
.
MultiKeyPollingIndex
)
}
func
TestSaveStatusStateFromSingleKeySnapshotPreservesUnownedColumns
(
t
*
testing
.
T
)
{
setupChannelStatusTest
(
t
)
channel
:=
Channel
{
Name
:
"single-key-status"
,
Key
:
"original-key"
,
Status
:
common
.
ChannelStatusEnabled
,
Models
:
"original-model"
,
Group
:
"default"
,
UsedQuota
:
100
,
ChannelInfo
:
ChannelInfo
{},
}
require
.
NoError
(
t
,
DB
.
Create
(
&
channel
)
.
Error
)
stale
,
err
:=
GetChannelById
(
channel
.
Id
,
true
)
require
.
NoError
(
t
,
err
)
concurrentChannelInfo
:=
ChannelInfo
{
IsMultiKey
:
true
,
MultiKeySize
:
2
,
MultiKeyMode
:
constant
.
MultiKeyModePolling
,
MultiKeyPollingIndex
:
1
,
}
require
.
NoError
(
t
,
DB
.
Model
(
&
Channel
{})
.
Where
(
"id = ?"
,
channel
.
Id
)
.
Updates
(
map
[
string
]
any
{
"key"
:
"rotated-key"
,
"used_quota"
:
gorm
.
Expr
(
"used_quota + ?"
,
250
),
"models"
:
"concurrent-model"
,
"channel_info"
:
concurrentChannelInfo
,
})
.
Error
)
stale
.
Status
=
common
.
ChannelStatusManuallyDisabled
stale
.
SetOtherInfo
(
map
[
string
]
interface
{}{
"status_reason"
:
"manual operation"
,
"status_time"
:
int64
(
1234
),
})
require
.
NoError
(
t
,
stale
.
saveStatusState
())
var
stored
Channel
require
.
NoError
(
t
,
DB
.
First
(
&
stored
,
channel
.
Id
)
.
Error
)
assert
.
Equal
(
t
,
common
.
ChannelStatusManuallyDisabled
,
stored
.
Status
)
assert
.
Equal
(
t
,
"rotated-key"
,
stored
.
Key
)
assert
.
Equal
(
t
,
int64
(
350
),
stored
.
UsedQuota
)
assert
.
Equal
(
t
,
"concurrent-model"
,
stored
.
Models
)
assert
.
Equal
(
t
,
concurrentChannelInfo
,
stored
.
ChannelInfo
)
otherInfo
:=
stored
.
GetOtherInfo
()
assert
.
Equal
(
t
,
"manual operation"
,
otherInfo
[
"status_reason"
])
assert
.
Equal
(
t
,
float64
(
1234
),
otherInfo
[
"status_time"
])
}
model/quota_reserve.go
View file @
ccd535ef
...
@@ -2,8 +2,11 @@ package model
...
@@ -2,8 +2,11 @@ package model
import
(
import
(
"context"
"context"
"errors"
"fmt"
"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/common"
"gorm.io/gorm"
)
)
type
cacheQuotaResult
int
type
cacheQuotaResult
int
...
@@ -14,11 +17,6 @@ const (
...
@@ -14,11 +17,6 @@ const (
cacheQuotaMiss
cacheQuotaMiss
)
)
// 下列脚本都是守卫式的:只在完整哈希(Id 匹配且配额字段存在)上操作,
// 哈希缺失时返回 miss 而不是创建残缺哈希。脚本不修改 TTL(HINCRBY 天然保留
// 水合时设置的 TTL),因此即使某个写库路径绕过了缓存,偏差也会在一个 TTL
// 窗口内随缓存过期而自愈。
const
userQuotaReserveScript
=
`
const
userQuotaReserveScript
=
`
if tonumber(redis.call('HGET', KEYS[1], 'Id') or '0') ~= tonumber(ARGV[2])
if tonumber(redis.call('HGET', KEYS[1], 'Id') or '0') ~= tonumber(ARGV[2])
or tonumber(redis.call('HGET', KEYS[1], 'CacheSchema') or '0') ~= tonumber(ARGV[3])
or tonumber(redis.call('HGET', KEYS[1], 'CacheSchema') or '0') ~= tonumber(ARGV[3])
...
@@ -41,6 +39,32 @@ end
...
@@ -41,6 +39,32 @@ end
redis.call('HINCRBY', KEYS[1], 'Quota', tonumber(ARGV[1]))
redis.call('HINCRBY', KEYS[1], 'Quota', tonumber(ARGV[1]))
return 1`
return 1`
const
tokenQuotaReserveScript
=
`
if tonumber(redis.call('HGET', KEYS[1], 'Id') or '0') ~= tonumber(ARGV[2])
or redis.call('HEXISTS', KEYS[1], 'RemainQuota') == 0
or redis.call('HEXISTS', KEYS[1], 'UsedQuota') == 0 then
return -1
end
local remain = tonumber(redis.call('HGET', KEYS[1], 'RemainQuota'))
if remain == nil or remain < tonumber(ARGV[1]) then
return 0
end
redis.call('HINCRBY', KEYS[1], 'RemainQuota', -tonumber(ARGV[1]))
redis.call('HINCRBY', KEYS[1], 'UsedQuota', tonumber(ARGV[1]))
redis.call('HSET', KEYS[1], 'AccessedTime', ARGV[3])
return 1`
const
tokenQuotaDeltaScript
=
`
if tonumber(redis.call('HGET', KEYS[1], 'Id') or '0') ~= tonumber(ARGV[2])
or redis.call('HEXISTS', KEYS[1], 'RemainQuota') == 0
or redis.call('HEXISTS', KEYS[1], 'UsedQuota') == 0 then
return -1
end
redis.call('HINCRBY', KEYS[1], 'RemainQuota', tonumber(ARGV[1]))
redis.call('HINCRBY', KEYS[1], 'UsedQuota', -tonumber(ARGV[1]))
redis.call('HSET', KEYS[1], 'AccessedTime', ARGV[3])
return 1`
func
quotaResultFromLua
(
result
int
,
err
error
)
(
cacheQuotaResult
,
error
)
{
func
quotaResultFromLua
(
result
int
,
err
error
)
(
cacheQuotaResult
,
error
)
{
if
err
!=
nil
{
if
err
!=
nil
{
return
cacheQuotaMiss
,
err
return
cacheQuotaMiss
,
err
...
@@ -66,3 +90,151 @@ func cacheApplyUserQuotaDelta(userID int, delta int64) (cacheQuotaResult, error)
...
@@ -66,3 +90,151 @@ func cacheApplyUserQuotaDelta(userID int, delta int64) (cacheQuotaResult, error)
[]
string
{
getUserCacheKey
(
userID
)},
delta
,
userID
,
userCacheSchemaVersion
)
.
Int
()
[]
string
{
getUserCacheKey
(
userID
)},
delta
,
userID
,
userCacheSchemaVersion
)
.
Int
()
return
quotaResultFromLua
(
result
,
err
)
return
quotaResultFromLua
(
result
,
err
)
}
}
func
cacheTryReserveTokenQuota
(
id
int
,
key
string
,
amount
int64
)
(
cacheQuotaResult
,
error
)
{
result
,
err
:=
common
.
RDB
.
Eval
(
context
.
Background
(),
tokenQuotaReserveScript
,
[]
string
{
getTokenCacheKey
(
key
)},
amount
,
id
,
common
.
GetTimestamp
())
.
Int
()
return
quotaResultFromLua
(
result
,
err
)
}
func
cacheApplyTokenQuotaDelta
(
id
int
,
key
string
,
delta
int64
)
(
cacheQuotaResult
,
error
)
{
result
,
err
:=
common
.
RDB
.
Eval
(
context
.
Background
(),
tokenQuotaDeltaScript
,
[]
string
{
getTokenCacheKey
(
key
)},
delta
,
id
,
common
.
GetTimestamp
())
.
Int
()
return
quotaResultFromLua
(
result
,
err
)
}
// persistUserQuotaDelta 把已在缓存侧预扣成功的增量落库;批量模式下入队,
// 直写模式下要求行存在(用户已删除时报错,交由调用方补偿缓存)。
func
persistUserQuotaDelta
(
id
int
,
delta
int
)
error
{
if
common
.
BatchUpdateEnabled
{
addNewRecord
(
BatchUpdateTypeUserQuota
,
id
,
delta
)
return
nil
}
result
:=
DB
.
Model
(
&
User
{})
.
Where
(
"id = ?"
,
id
)
.
Update
(
"quota"
,
gorm
.
Expr
(
"quota + ?"
,
delta
))
if
result
.
Error
!=
nil
{
return
result
.
Error
}
if
result
.
RowsAffected
!=
1
{
return
gorm
.
ErrRecordNotFound
}
return
nil
}
func
persistTokenQuotaDelta
(
id
int
,
delta
int
)
error
{
if
common
.
BatchUpdateEnabled
{
addNewRecord
(
BatchUpdateTypeTokenQuota
,
id
,
delta
)
return
nil
}
result
:=
DB
.
Model
(
&
Token
{})
.
Where
(
"id = ?"
,
id
)
.
Updates
(
map
[
string
]
interface
{}{
"remain_quota"
:
gorm
.
Expr
(
"remain_quota + ?"
,
delta
),
"used_quota"
:
gorm
.
Expr
(
"used_quota - ?"
,
delta
),
"accessed_time"
:
common
.
GetTimestamp
(),
},
)
if
result
.
Error
!=
nil
{
return
result
.
Error
}
if
result
.
RowsAffected
!=
1
{
return
gorm
.
ErrRecordNotFound
}
return
nil
}
func
reserveUserQuotaDB
(
id
int
,
quota
int
)
(
bool
,
error
)
{
result
:=
DB
.
Model
(
&
User
{})
.
Where
(
"id = ? AND quota >= ?"
,
id
,
quota
)
.
Update
(
"quota"
,
gorm
.
Expr
(
"quota - ?"
,
quota
))
return
result
.
RowsAffected
==
1
,
result
.
Error
}
func
reserveTokenQuotaDB
(
id
int
,
quota
int
)
(
bool
,
error
)
{
result
:=
DB
.
Model
(
&
Token
{})
.
Where
(
"id = ? AND remain_quota >= ?"
,
id
,
quota
)
.
Updates
(
map
[
string
]
interface
{}{
"remain_quota"
:
gorm
.
Expr
(
"remain_quota - ?"
,
quota
),
"used_quota"
:
gorm
.
Expr
(
"used_quota + ?"
,
quota
),
"accessed_time"
:
common
.
GetTimestamp
(),
})
return
result
.
RowsAffected
==
1
,
result
.
Error
}
// TryReserveUserQuota atomically checks and deducts a user's wallet quota.
// 缓存命中时以缓存余额为准(避免批量模式下过期的数据库余额放大并发超扣);
// Redis 异常或水合失败时降级为数据库条件更新,保证服务可用。
func
TryReserveUserQuota
(
id
int
,
quota
int
)
(
bool
,
error
)
{
if
quota
<
0
{
return
false
,
errors
.
New
(
"quota 不能为负数!"
)
}
if
quota
==
0
{
return
true
,
nil
}
if
!
common
.
RedisEnabled
{
return
reserveUserQuotaDB
(
id
,
quota
)
}
result
,
err
:=
cacheTryReserveUserQuota
(
id
,
int64
(
quota
))
if
err
==
nil
&&
result
==
cacheQuotaMiss
{
if
_
,
hydrateErr
:=
GetUserCache
(
id
);
hydrateErr
==
nil
{
result
,
err
=
cacheTryReserveUserQuota
(
id
,
int64
(
quota
))
}
}
if
err
!=
nil
||
result
==
cacheQuotaMiss
{
if
err
!=
nil
{
common
.
SysLog
(
"user quota cache reserve unavailable, falling back to database: "
+
err
.
Error
())
}
return
reserveUserQuotaDB
(
id
,
quota
)
}
if
result
==
cacheQuotaInsufficient
{
return
false
,
nil
}
if
err
=
persistUserQuotaDelta
(
id
,
-
quota
);
err
!=
nil
{
compensated
,
compensateErr
:=
cacheApplyUserQuotaDelta
(
id
,
int64
(
quota
))
if
compensateErr
!=
nil
||
compensated
!=
cacheQuotaOK
{
common
.
SysError
(
fmt
.
Sprintf
(
"failed to compensate reserved user quota: result=%d error=%v"
,
compensated
,
compensateErr
))
}
return
false
,
err
}
return
true
,
nil
}
// TryReserveTokenQuota atomically checks and deducts a token quota. Unlimited
// tokens skip the balance check but still update remain/used accounting.
func
TryReserveTokenQuota
(
id
int
,
key
string
,
quota
int
,
unlimited
bool
)
(
bool
,
error
)
{
if
quota
<
0
{
return
false
,
errors
.
New
(
"quota 不能为负数!"
)
}
if
quota
==
0
{
return
true
,
nil
}
if
unlimited
{
return
true
,
DecreaseTokenQuota
(
id
,
key
,
quota
)
}
if
!
common
.
RedisEnabled
{
return
reserveTokenQuotaDB
(
id
,
quota
)
}
result
,
err
:=
cacheTryReserveTokenQuota
(
id
,
key
,
int64
(
quota
))
if
err
==
nil
&&
result
==
cacheQuotaMiss
{
if
_
,
hydrateErr
:=
GetTokenByKey
(
key
,
true
);
hydrateErr
==
nil
{
result
,
err
=
cacheTryReserveTokenQuota
(
id
,
key
,
int64
(
quota
))
}
}
if
err
!=
nil
||
result
==
cacheQuotaMiss
{
if
err
!=
nil
{
common
.
SysLog
(
"token quota cache reserve unavailable, falling back to database: "
+
err
.
Error
())
}
return
reserveTokenQuotaDB
(
id
,
quota
)
}
if
result
==
cacheQuotaInsufficient
{
return
false
,
nil
}
if
err
=
persistTokenQuotaDelta
(
id
,
-
quota
);
err
!=
nil
{
compensated
,
compensateErr
:=
cacheApplyTokenQuotaDelta
(
id
,
key
,
int64
(
quota
))
if
compensateErr
!=
nil
||
compensated
!=
cacheQuotaOK
{
common
.
SysError
(
fmt
.
Sprintf
(
"failed to compensate reserved token quota: result=%d error=%v"
,
compensated
,
compensateErr
))
}
return
false
,
err
}
return
true
,
nil
}
model/quota_reserve_test.go
0 → 100644
View file @
ccd535ef
package
model
import
(
"testing"
"time"
"github.com/QuantumNous/new-api/common"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func
createReserveTestUser
(
t
*
testing
.
T
,
quota
int
)
User
{
t
.
Helper
()
user
:=
User
{
Username
:
"reserve-user-"
+
common
.
GetRandomString
(
6
),
Password
:
"unused-password-hash"
,
Role
:
common
.
RoleCommonUser
,
Status
:
common
.
UserStatusEnabled
,
Group
:
"default"
,
AuthVersion
:
1
,
Quota
:
quota
,
AffCode
:
"reserve-aff-"
+
common
.
GetRandomString
(
8
),
}
require
.
NoError
(
t
,
DB
.
Create
(
&
user
)
.
Error
)
return
user
}
func
createReserveTestToken
(
t
*
testing
.
T
,
remainQuota
int
)
Token
{
t
.
Helper
()
token
:=
Token
{
UserId
:
1
,
Key
:
"reserve-token-"
+
common
.
GetRandomString
(
8
),
Name
:
"reserve-test"
,
Status
:
common
.
TokenStatusEnabled
,
ExpiredTime
:
-
1
,
RemainQuota
:
remainQuota
,
}
require
.
NoError
(
t
,
token
.
Insert
())
return
token
}
func
getUserQuotaFromDB
(
t
*
testing
.
T
,
id
int
)
int
{
t
.
Helper
()
var
user
User
require
.
NoError
(
t
,
DB
.
Select
(
"quota"
)
.
First
(
&
user
,
id
)
.
Error
)
return
user
.
Quota
}
func
getTokenFromDB
(
t
*
testing
.
T
,
id
int
)
Token
{
t
.
Helper
()
var
token
Token
require
.
NoError
(
t
,
DB
.
First
(
&
token
,
id
)
.
Error
)
return
token
}
func
resetBatchUpdateTestState
(
t
*
testing
.
T
)
{
t
.
Helper
()
oldBatchEnabled
:=
common
.
BatchUpdateEnabled
common
.
BatchUpdateEnabled
=
false
for
i
:=
0
;
i
<
BatchUpdateTypeCount
;
i
++
{
batchUpdateLocks
[
i
]
.
Lock
()
batchUpdateStores
[
i
]
=
make
(
map
[
int
]
int
)
batchUpdateLocks
[
i
]
.
Unlock
()
}
t
.
Cleanup
(
func
()
{
common
.
BatchUpdateEnabled
=
oldBatchEnabled
for
i
:=
0
;
i
<
BatchUpdateTypeCount
;
i
++
{
batchUpdateLocks
[
i
]
.
Lock
()
batchUpdateStores
[
i
]
=
make
(
map
[
int
]
int
)
batchUpdateLocks
[
i
]
.
Unlock
()
}
})
}
func
TestTryReserveQuotaWithoutRedis
(
t
*
testing
.
T
)
{
truncateTables
(
t
)
resetBatchUpdateTestState
(
t
)
user
:=
createReserveTestUser
(
t
,
100
)
reserved
,
err
:=
TryReserveUserQuota
(
user
.
Id
,
60
)
require
.
NoError
(
t
,
err
)
assert
.
True
(
t
,
reserved
)
assert
.
Equal
(
t
,
40
,
getUserQuotaFromDB
(
t
,
user
.
Id
))
reserved
,
err
=
TryReserveUserQuota
(
user
.
Id
,
41
)
require
.
NoError
(
t
,
err
)
assert
.
False
(
t
,
reserved
)
assert
.
Equal
(
t
,
40
,
getUserQuotaFromDB
(
t
,
user
.
Id
))
token
:=
createReserveTestToken
(
t
,
80
)
reserved
,
err
=
TryReserveTokenQuota
(
token
.
Id
,
token
.
Key
,
25
,
false
)
require
.
NoError
(
t
,
err
)
assert
.
True
(
t
,
reserved
)
reloaded
:=
getTokenFromDB
(
t
,
token
.
Id
)
assert
.
Equal
(
t
,
55
,
reloaded
.
RemainQuota
)
assert
.
Equal
(
t
,
25
,
reloaded
.
UsedQuota
)
reserved
,
err
=
TryReserveTokenQuota
(
token
.
Id
,
token
.
Key
,
56
,
false
)
require
.
NoError
(
t
,
err
)
assert
.
False
(
t
,
reserved
)
assert
.
Equal
(
t
,
55
,
getTokenFromDB
(
t
,
token
.
Id
)
.
RemainQuota
)
}
func
TestRedisBatchReserveNeverFallsBackToStaleDatabaseBalance
(
t
*
testing
.
T
)
{
truncateTables
(
t
)
resetBatchUpdateTestState
(
t
)
useUserCacheMiniRedis
(
t
)
common
.
BatchUpdateEnabled
=
true
user
:=
createReserveTestUser
(
t
,
10
)
reserved
,
err
:=
TryReserveUserQuota
(
user
.
Id
,
8
)
require
.
NoError
(
t
,
err
)
assert
.
True
(
t
,
reserved
)
assert
.
Equal
(
t
,
10
,
getUserQuotaFromDB
(
t
,
user
.
Id
),
"batch delta is not flushed yet"
)
reserved
,
err
=
TryReserveUserQuota
(
user
.
Id
,
3
)
require
.
NoError
(
t
,
err
)
assert
.
False
(
t
,
reserved
,
"stale DB balance must not authorize a second spend"
)
cachedUser
,
err
:=
GetUserCache
(
user
.
Id
)
require
.
NoError
(
t
,
err
)
assert
.
Equal
(
t
,
2
,
cachedUser
.
Quota
)
token
:=
createReserveTestToken
(
t
,
9
)
reserved
,
err
=
TryReserveTokenQuota
(
token
.
Id
,
token
.
Key
,
7
,
false
)
require
.
NoError
(
t
,
err
)
assert
.
True
(
t
,
reserved
)
reserved
,
err
=
TryReserveTokenQuota
(
token
.
Id
,
token
.
Key
,
3
,
false
)
require
.
NoError
(
t
,
err
)
assert
.
False
(
t
,
reserved
)
assert
.
Equal
(
t
,
9
,
getTokenFromDB
(
t
,
token
.
Id
)
.
RemainQuota
)
batchUpdate
()
assert
.
Equal
(
t
,
2
,
getUserQuotaFromDB
(
t
,
user
.
Id
))
reloadedToken
:=
getTokenFromDB
(
t
,
token
.
Id
)
assert
.
Equal
(
t
,
2
,
reloadedToken
.
RemainQuota
)
assert
.
Equal
(
t
,
7
,
reloadedToken
.
UsedQuota
)
}
func
TestReserveFallsBackToDatabaseWhenRedisIsUnavailable
(
t
*
testing
.
T
)
{
truncateTables
(
t
)
resetBatchUpdateTestState
(
t
)
server
:=
useUserCacheMiniRedis
(
t
)
user
:=
createReserveTestUser
(
t
,
20
)
require
.
NoError
(
t
,
populateUserCache
(
user
))
server
.
Close
()
// Redis 故障时降级为数据库条件更新:服务保持可用且不会超扣。
reserved
,
err
:=
TryReserveUserQuota
(
user
.
Id
,
5
)
require
.
NoError
(
t
,
err
)
assert
.
True
(
t
,
reserved
)
assert
.
Equal
(
t
,
15
,
getUserQuotaFromDB
(
t
,
user
.
Id
))
reserved
,
err
=
TryReserveUserQuota
(
user
.
Id
,
16
)
require
.
NoError
(
t
,
err
)
assert
.
False
(
t
,
reserved
)
assert
.
Equal
(
t
,
15
,
getUserQuotaFromDB
(
t
,
user
.
Id
))
}
func
TestSynchronousReserveCompensatesCacheWhenPersistenceFails
(
t
*
testing
.
T
)
{
truncateTables
(
t
)
resetBatchUpdateTestState
(
t
)
useUserCacheMiniRedis
(
t
)
user
:=
createReserveTestUser
(
t
,
10
)
require
.
NoError
(
t
,
populateUserCache
(
user
))
require
.
NoError
(
t
,
DB
.
Delete
(
&
user
)
.
Error
)
reserved
,
err
:=
TryReserveUserQuota
(
user
.
Id
,
6
)
assert
.
False
(
t
,
reserved
)
assert
.
ErrorIs
(
t
,
err
,
gorm
.
ErrRecordNotFound
)
cached
,
cacheErr
:=
cacheGetUserBase
(
user
.
Id
)
require
.
NoError
(
t
,
cacheErr
)
assert
.
Equal
(
t
,
10
,
cached
.
Quota
)
token
:=
createReserveTestToken
(
t
,
12
)
_
,
err
=
GetTokenByKey
(
token
.
Key
,
true
)
require
.
NoError
(
t
,
err
)
require
.
NoError
(
t
,
DB
.
Delete
(
&
token
)
.
Error
)
reserved
,
err
=
TryReserveTokenQuota
(
token
.
Id
,
token
.
Key
,
7
,
false
)
assert
.
False
(
t
,
reserved
)
assert
.
ErrorIs
(
t
,
err
,
gorm
.
ErrRecordNotFound
)
cachedToken
,
cacheErr
:=
cacheGetTokenByKey
(
token
.
Key
)
require
.
NoError
(
t
,
cacheErr
)
assert
.
Equal
(
t
,
12
,
cachedToken
.
RemainQuota
)
assert
.
Zero
(
t
,
cachedToken
.
UsedQuota
)
}
func
TestTokenCacheInitPreservesLiveQuotaAndFenceBlocksStaleSnapshot
(
t
*
testing
.
T
)
{
truncateTables
(
t
)
resetBatchUpdateTestState
(
t
)
server
:=
useUserCacheMiniRedis
(
t
)
token
:=
createReserveTestToken
(
t
,
100
)
loaded
,
err
:=
GetTokenByKey
(
token
.
Key
,
true
)
require
.
NoError
(
t
,
err
)
stale
:=
*
loaded
result
,
err
:=
cacheApplyTokenQuotaDelta
(
token
.
Id
,
token
.
Key
,
-
70
)
require
.
NoError
(
t
,
err
)
require
.
Equal
(
t
,
cacheQuotaOK
,
result
)
// 已存在的哈希只刷新 TTL:数据库快照不得覆盖已被原子预扣的余额。
code
,
err
:=
cacheInitToken
(
stale
)
require
.
NoError
(
t
,
err
)
assert
.
Equal
(
t
,
2
,
code
)
cached
,
err
:=
cacheGetTokenByKey
(
token
.
Key
)
require
.
NoError
(
t
,
err
)
assert
.
Equal
(
t
,
30
,
cached
.
RemainQuota
)
// 变更期间:fence 删除缓存并拦截并发读者手中的过期快照。
require
.
NoError
(
t
,
invalidateTokenCacheForMutation
(
token
.
Key
))
code
,
err
=
cacheInitToken
(
stale
)
require
.
NoError
(
t
,
err
)
assert
.
Zero
(
t
,
code
,
"the pre-mutation snapshot must not be published while fenced"
)
_
,
err
=
cacheGetTokenByKey
(
token
.
Key
)
assert
.
Error
(
t
,
err
)
// fence 过期后可重新从数据库水合。
server
.
FastForward
(
time
.
Duration
(
tokenCacheFenceSeconds
+
1
)
*
time
.
Second
)
fresh
,
err
:=
GetTokenByKey
(
token
.
Key
,
false
)
require
.
NoError
(
t
,
err
)
assert
.
Equal
(
t
,
100
,
fresh
.
RemainQuota
)
cached
,
err
=
cacheGetTokenByKey
(
token
.
Key
)
require
.
NoError
(
t
,
err
)
assert
.
Equal
(
t
,
100
,
cached
.
RemainQuota
)
}
model/subscription.go
View file @
ccd535ef
...
@@ -600,6 +600,12 @@ func CompleteSubscriptionOrder(tradeNo string, providerPayload string, expectedP
...
@@ -600,6 +600,12 @@ func CompleteSubscriptionOrder(tradeNo string, providerPayload string, expectedP
if
!
plan
.
Enabled
{
if
!
plan
.
Enabled
{
// still allow completion for already purchased orders
// still allow completion for already purchased orders
}
}
// 锁定用户行:并发完成同一用户的不同订单(包括多实例部署下)时,
// 使 CreateUserSubscriptionFromPlanTx 的 MaxPurchasePerUser 检查按用户串行。
var
userRow
User
if
err
:=
lockForUpdate
(
tx
)
.
Select
(
"id"
)
.
Where
(
"id = ?"
,
order
.
UserId
)
.
First
(
&
userRow
)
.
Error
;
err
!=
nil
{
return
err
}
subscription
,
err
:=
CreateUserSubscriptionFromPlanTx
(
tx
,
order
.
UserId
,
plan
,
"order"
)
subscription
,
err
:=
CreateUserSubscriptionFromPlanTx
(
tx
,
order
.
UserId
,
plan
,
"order"
)
if
err
!=
nil
{
if
err
!=
nil
{
return
err
return
err
...
@@ -712,6 +718,11 @@ func AdminBindSubscription(userId int, planId int, sourceNote string) (string, e
...
@@ -712,6 +718,11 @@ func AdminBindSubscription(userId int, planId int, sourceNote string) (string, e
}
}
groupChanged
:=
false
groupChanged
:=
false
err
=
DB
.
Transaction
(
func
(
tx
*
gorm
.
DB
)
error
{
err
=
DB
.
Transaction
(
func
(
tx
*
gorm
.
DB
)
error
{
// 与 CompleteSubscriptionOrder 一致:先锁用户行,再做购买次数检查。
var
userRow
User
if
err
:=
lockForUpdate
(
tx
)
.
Select
(
"id"
)
.
Where
(
"id = ?"
,
userId
)
.
First
(
&
userRow
)
.
Error
;
err
!=
nil
{
return
err
}
subscription
,
err
:=
CreateUserSubscriptionFromPlanTx
(
tx
,
userId
,
plan
,
"admin"
)
subscription
,
err
:=
CreateUserSubscriptionFromPlanTx
(
tx
,
userId
,
plan
,
"admin"
)
if
err
==
nil
{
if
err
==
nil
{
groupChanged
=
subscription
.
PrevUserGroup
!=
""
groupChanged
=
subscription
.
PrevUserGroup
!=
""
...
@@ -737,9 +748,8 @@ func calcSubscriptionBalanceQuota(priceAmount float64) (int, error) {
...
@@ -737,9 +748,8 @@ func calcSubscriptionBalanceQuota(priceAmount float64) (int, error) {
}
}
quota
:=
decimal
.
NewFromFloat
(
priceAmount
)
.
quota
:=
decimal
.
NewFromFloat
(
priceAmount
)
.
Mul
(
decimal
.
NewFromFloat
(
common
.
QuotaPerUnit
))
.
Mul
(
decimal
.
NewFromFloat
(
common
.
QuotaPerUnit
))
.
Ceil
()
.
Ceil
()
IntPart
()
return
common
.
QuotaFromDecimalStrict
(
quota
)
return
int
(
quota
),
nil
}
}
// PurchaseSubscriptionWithBalance creates a subscription by deducting the user's wallet quota.
// PurchaseSubscriptionWithBalance creates a subscription by deducting the user's wallet quota.
...
...
model/token.go
View file @
ccd535ef
...
@@ -274,27 +274,10 @@ func GetTokenById(id int) (*Token, error) {
...
@@ -274,27 +274,10 @@ func GetTokenById(id int) (*Token, error) {
token
:=
Token
{
Id
:
id
}
token
:=
Token
{
Id
:
id
}
var
err
error
=
nil
var
err
error
=
nil
err
=
DB
.
First
(
&
token
,
"id = ?"
,
id
)
.
Error
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
return
&
token
,
err
}
}
func
GetTokenByKey
(
key
string
,
fromDB
bool
)
(
token
*
Token
,
err
error
)
{
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
{
if
!
fromDB
&&
common
.
RedisEnabled
{
// Try Redis first
// Try Redis first
token
,
err
:=
cacheGetTokenByKey
(
key
)
token
,
err
:=
cacheGetTokenByKey
(
key
)
...
@@ -303,9 +286,18 @@ func GetTokenByKey(key string, fromDB bool) (token *Token, err error) {
...
@@ -303,9 +286,18 @@ func GetTokenByKey(key string, fromDB bool) (token *Token, err error) {
}
}
// Don't return error - fall through to DB
// Don't return error - fall through to DB
}
}
fromDB
=
true
token
=
&
Token
{}
err
=
DB
.
Where
(
commonKeyCol
+
" = ?"
,
key
)
.
First
(
&
token
)
.
Error
if
err
=
DB
.
Where
(
commonKeyCol
+
" = ?"
,
key
)
.
First
(
token
)
.
Error
;
err
!=
nil
{
return
token
,
err
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
{
func
(
token
*
Token
)
Insert
()
error
{
...
@@ -316,47 +308,27 @@ 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
// Update Make sure your token's fields is completed, because this will update non-zero values
func
(
token
*
Token
)
Update
()
(
err
error
)
{
func
(
token
*
Token
)
Update
()
(
err
error
)
{
err
=
DB
.
Model
(
token
)
.
Select
(
"name"
,
"status"
,
"expired_time"
,
"remain_quota"
,
"unlimited_quota"
,
// 写库前失效缓存并设置 fence,防止并发读者把过期快照重新写回缓存。
"model_limits_enabled"
,
"model_limits"
,
"allow_ips"
,
"group"
,
"cross_group_retry"
,
"auto_groups"
)
.
Updates
(
token
)
.
Error
if
cacheErr
:=
invalidateTokenCacheForMutation
(
token
.
Key
);
cacheErr
!=
nil
{
if
shouldUpdateRedis
(
true
,
err
)
{
common
.
SysLog
(
"failed to invalidate token cache before update: "
+
cacheErr
.
Error
())
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
())
}
}
}
}
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
)
{
func
(
token
*
Token
)
SelectUpdate
()
(
err
error
)
{
defer
func
()
{
if
cacheErr
:=
invalidateTokenCacheForMutation
(
token
.
Key
);
cacheErr
!=
nil
{
if
shouldUpdateRedis
(
true
,
err
)
{
common
.
SysLog
(
"failed to invalidate token cache before status update: "
+
cacheErr
.
Error
())
gopool
.
Go
(
func
()
{
}
err
:=
cacheSetToken
(
*
token
)
if
err
!=
nil
{
common
.
SysLog
(
"failed to update token cache: "
+
err
.
Error
())
}
})
}
}()
// This can update zero values
// This can update zero values
return
DB
.
Model
(
token
)
.
Select
(
"accessed_time"
,
"status"
)
.
Updates
(
token
)
.
Error
return
DB
.
Model
(
token
)
.
Select
(
"accessed_time"
,
"status"
)
.
Updates
(
token
)
.
Error
}
}
func
(
token
*
Token
)
Delete
()
(
err
error
)
{
func
(
token
*
Token
)
Delete
()
(
err
error
)
{
defer
func
()
{
if
cacheErr
:=
invalidateTokenCacheForMutation
(
token
.
Key
);
cacheErr
!=
nil
{
if
shouldUpdateRedis
(
true
,
err
)
{
common
.
SysLog
(
"failed to invalidate token cache before delete: "
+
cacheErr
.
Error
())
gopool
.
Go
(
func
()
{
}
err
:=
cacheDeleteToken
(
token
.
Key
)
return
DB
.
Delete
(
token
)
.
Error
if
err
!=
nil
{
common
.
SysLog
(
"failed to delete token cache: "
+
err
.
Error
())
}
})
}
}()
err
=
DB
.
Delete
(
token
)
.
Error
return
err
}
}
func
(
token
*
Token
)
IsModelLimitsEnabled
()
bool
{
func
(
token
*
Token
)
IsModelLimitsEnabled
()
bool
{
...
@@ -408,8 +380,9 @@ func IncreaseTokenQuota(tokenId int, key string, quota int) (err error) {
...
@@ -408,8 +380,9 @@ func IncreaseTokenQuota(tokenId int, key string, quota int) (err error) {
}
}
if
common
.
RedisEnabled
{
if
common
.
RedisEnabled
{
gopool
.
Go
(
func
()
{
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
())
common
.
SysLog
(
"failed to increase token quota: "
+
err
.
Error
())
}
}
})
})
...
@@ -438,8 +411,7 @@ func DecreaseTokenQuota(id int, key string, quota int) (err error) {
...
@@ -438,8 +411,7 @@ func DecreaseTokenQuota(id int, key string, quota int) (err error) {
}
}
if
common
.
RedisEnabled
{
if
common
.
RedisEnabled
{
gopool
.
Go
(
func
()
{
gopool
.
Go
(
func
()
{
err
:=
cacheDecrTokenQuota
(
key
,
int64
(
quota
))
if
_
,
err
:=
cacheApplyTokenQuotaDelta
(
id
,
key
,
int64
(
-
quota
));
err
!=
nil
{
if
err
!=
nil
{
common
.
SysLog
(
"failed to decrease token quota: "
+
err
.
Error
())
common
.
SysLog
(
"failed to decrease token quota: "
+
err
.
Error
())
}
}
})
})
...
@@ -482,6 +454,9 @@ func BatchDeleteTokens(ids []int, userId int) (int, error) {
...
@@ -482,6 +454,9 @@ func BatchDeleteTokens(ids []int, userId int) (int, error) {
tx
.
Rollback
()
tx
.
Rollback
()
return
0
,
err
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
{
if
err
:=
tx
.
Where
(
"user_id = ? AND id IN (?)"
,
userId
,
ids
)
.
Delete
(
&
Token
{})
.
Error
;
err
!=
nil
{
tx
.
Rollback
()
tx
.
Rollback
()
...
@@ -492,14 +467,6 @@ func BatchDeleteTokens(ids []int, userId int) (int, error) {
...
@@ -492,14 +467,6 @@ func BatchDeleteTokens(ids []int, userId int) (int, error) {
return
0
,
err
return
0
,
err
}
}
if
common
.
RedisEnabled
{
gopool
.
Go
(
func
()
{
for
_
,
t
:=
range
tokens
{
_
=
cacheDeleteToken
(
t
.
Key
)
}
})
}
return
len
(
tokens
),
nil
return
len
(
tokens
),
nil
}
}
...
@@ -540,7 +507,7 @@ func invalidateTokensCache(tokens []Token) error {
...
@@ -540,7 +507,7 @@ func invalidateTokensCache(tokens []Token) error {
if
t
.
Key
==
""
{
if
t
.
Key
==
""
{
continue
continue
}
}
if
err
:=
cacheDeleteToke
n
(
t
.
Key
);
err
!=
nil
&&
firstErr
==
nil
{
if
err
:=
invalidateTokenCacheForMutatio
n
(
t
.
Key
);
err
!=
nil
&&
firstErr
==
nil
{
firstErr
=
err
firstErr
=
err
}
}
}
}
...
...
model/token_auto_groups_cache_test.go
View file @
ccd535ef
...
@@ -19,7 +19,7 @@ func TestTokenAutoGroupsRoundTripThroughRedisHashCache(t *testing.T) {
...
@@ -19,7 +19,7 @@ func TestTokenAutoGroupsRoundTripThroughRedisHashCache(t *testing.T) {
AutoGroups
:
`["vip","default"]`
,
AutoGroups
:
`["vip","default"]`
,
}
}
require
.
NoError
(
t
,
cacheSetToken
(
token
))
require
.
NoError
(
t
,
cacheSetToken
ForTest
(
token
))
cached
,
err
:=
cacheGetTokenByKey
(
token
.
Key
)
cached
,
err
:=
cacheGetTokenByKey
(
token
.
Key
)
require
.
NoError
(
t
,
err
)
require
.
NoError
(
t
,
err
)
assert
.
Equal
(
t
,
token
.
AutoGroups
,
cached
.
AutoGroups
)
assert
.
Equal
(
t
,
token
.
AutoGroups
,
cached
.
AutoGroups
)
...
@@ -43,7 +43,7 @@ func TestTokenUpdateSynchronouslyNarrowsPreheatedAutoGroupsCache(t *testing.T) {
...
@@ -43,7 +43,7 @@ func TestTokenUpdateSynchronouslyNarrowsPreheatedAutoGroupsCache(t *testing.T) {
AutoGroups
:
`["default","vip"]`
,
AutoGroups
:
`["default","vip"]`
,
}
}
require
.
NoError
(
t
,
token
.
Insert
())
require
.
NoError
(
t
,
token
.
Insert
())
require
.
NoError
(
t
,
cacheSetToken
(
token
))
require
.
NoError
(
t
,
cacheSetToken
ForTest
(
token
))
preheated
,
err
:=
cacheGetTokenByKey
(
token
.
Key
)
preheated
,
err
:=
cacheGetTokenByKey
(
token
.
Key
)
require
.
NoError
(
t
,
err
)
require
.
NoError
(
t
,
err
)
...
@@ -51,7 +51,18 @@ func TestTokenUpdateSynchronouslyNarrowsPreheatedAutoGroupsCache(t *testing.T) {
...
@@ -51,7 +51,18 @@ func TestTokenUpdateSynchronouslyNarrowsPreheatedAutoGroupsCache(t *testing.T) {
require
.
NoError
(
t
,
token
.
SetAutoGroups
([]
string
{
"vip"
}))
require
.
NoError
(
t
,
token
.
SetAutoGroups
([]
string
{
"vip"
}))
require
.
NoError
(
t
,
token
.
Update
())
require
.
NoError
(
t
,
token
.
Update
())
immediate
,
err
:=
cacheGetTokenByKey
(
token
.
Key
)
// Update 是限制性变更:写库前删除缓存并设置 fence。缓存不再提供旧的
// 宽分组值,下一次读取必须看到收紧后的分组。
_
,
cacheErr
:=
cacheGetTokenByKey
(
token
.
Key
)
require
.
Error
(
t
,
cacheErr
,
"the pre-update cache entry must be invalidated"
)
reloaded
,
err
:=
GetTokenByKey
(
token
.
Key
,
false
)
require
.
NoError
(
t
,
err
)
require
.
NoError
(
t
,
err
)
assert
.
JSONEq
(
t
,
`["vip"]`
,
immediate
.
AutoGroups
)
assert
.
JSONEq
(
t
,
`["vip"]`
,
reloaded
.
AutoGroups
)
}
// cacheSetTokenForTest 以测试身份写入完整 token 缓存(含额度字段),
// 模拟“已水合”的缓存状态。
func
cacheSetTokenForTest
(
token
Token
)
error
{
_
,
err
:=
cacheInitToken
(
token
)
return
err
}
}
model/token_cache.go
View file @
ccd535ef
package
model
package
model
import
(
import
(
"context"
"fmt"
"fmt"
"strconv"
"time"
"time"
"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/constant"
)
)
func
cacheSetToken
(
token
Token
)
error
{
func
getTokenCacheKey
(
key
string
)
string
{
key
:=
common
.
GenerateHMAC
(
token
.
Key
)
return
fmt
.
Sprintf
(
"token:%s"
,
common
.
GenerateHMAC
(
key
))
token
.
Clean
()
err
:=
common
.
RedisHSetObj
(
fmt
.
Sprintf
(
"token:%s"
,
key
),
&
token
,
time
.
Duration
(
common
.
RedisKeyCacheSeconds
())
*
time
.
Second
)
if
err
!=
nil
{
return
err
}
return
nil
}
}
func
cacheDeleteToken
(
key
string
)
error
{
func
getTokenCacheFenceKey
(
key
string
)
string
{
key
=
common
.
GenerateHMAC
(
key
)
return
fmt
.
Sprintf
(
"token:fence:%s"
,
common
.
GenerateHMAC
(
key
))
err
:=
common
.
RedisDelKey
(
fmt
.
Sprintf
(
"token:%s"
,
key
))
if
err
!=
nil
{
return
err
}
return
nil
}
}
func
cacheIncrTokenQuota
(
key
string
,
increment
int64
)
error
{
func
tokenCacheTTLSeconds
()
int
{
key
=
common
.
GenerateHMAC
(
key
)
ttl
:=
common
.
RedisKeyCacheSeconds
()
err
:=
common
.
RedisHIncrBy
(
fmt
.
Sprintf
(
"token:%s"
,
key
),
constant
.
TokenFiledRemainQuota
,
increment
)
if
ttl
<=
0
{
if
err
!=
nil
{
return
60
return
err
}
}
return
ni
l
return
tt
l
}
}
func
cacheDecrTokenQuota
(
key
string
,
decrement
int64
)
error
{
// tokenCacheFenceSeconds must outlive a token mutation's database write plus
return
cacheIncrTokenQuota
(
key
,
-
decrement
)
// any in-flight reader's DB-read-to-cache-init gap. The fence is not deleted
}
// after commit; it expires naturally so a reader holding a pre-mutation
// snapshot cannot publish it right after the mutation cleared the cache.
// While the fence exists readers simply serve the database without caching.
const
tokenCacheFenceSeconds
=
10
func
cacheSetTokenField
(
key
string
,
field
string
,
value
string
)
error
{
// invalidateTokenCacheForMutation is called before a token metadata mutation
key
=
common
.
GenerateHMAC
(
key
)
// writes to the database: it raises the fence and drops the cached hash so no
err
:=
common
.
RedisHSetField
(
fmt
.
Sprintf
(
"token:%s"
,
key
),
field
,
value
)
// reader can act on (or re-publish) the pre-mutation state.
func
invalidateTokenCacheForMutation
(
key
string
)
error
{
if
!
common
.
RedisEnabled
||
key
==
""
{
return
nil
}
ctx
:=
context
.
Background
()
err
:=
common
.
RDB
.
Set
(
ctx
,
getTokenCacheFenceKey
(
key
),
1
,
time
.
Duration
(
tokenCacheFenceSeconds
)
*
time
.
Second
)
.
Err
()
if
err
!=
nil
{
if
err
!=
nil
{
return
err
return
err
}
}
return
nil
return
common
.
RDB
.
Del
(
ctx
,
getTokenCacheKey
(
key
))
.
Err
()
}
}
// CacheGetTokenByKey 从缓存中获取 token,如果缓存中不存在,则从数据库中获取
// cacheInitToken publishes a database snapshot only when no mutation fence is
// active and the hash is cold. An existing hash only gets its TTL refreshed:
// its RemainQuota may already be ahead of this snapshot because atomic
// pre-consume decrements Redis first, so a snapshot must never overwrite any
// field of a live hash.
// 返回值:0=被 fence 拦截,1=完成初始化,2=哈希已存在,仅刷新 TTL。
func
cacheInitToken
(
token
Token
)
(
int
,
error
)
{
if
!
common
.
RedisEnabled
{
return
0
,
nil
}
allowIps
:=
""
if
token
.
AllowIps
!=
nil
{
allowIps
=
*
token
.
AllowIps
}
const
script
=
`
if redis.call('EXISTS', KEYS[2]) == 1 then
return 0
end
if redis.call('EXISTS', KEYS[1]) == 1 then
redis.call('EXPIRE', KEYS[1], ARGV[17])
return 2
end
redis.call('HSET', KEYS[1],
'Id', ARGV[1], 'UserId', ARGV[2], 'Status', ARGV[3], 'Name', ARGV[4],
'CreatedTime', ARGV[5], 'AccessedTime', ARGV[6], 'ExpiredTime', ARGV[7],
'UnlimitedQuota', ARGV[8], 'ModelLimitsEnabled', ARGV[9], 'ModelLimits', ARGV[10],
'AllowIps', ARGV[11], 'Group', ARGV[12], 'CrossGroupRetry', ARGV[13],
'AutoGroups', ARGV[14], 'RemainQuota', ARGV[15], 'UsedQuota', ARGV[16])
redis.call('EXPIRE', KEYS[1], ARGV[17])
return 1`
return
common
.
RDB
.
Eval
(
context
.
Background
(),
script
,
[]
string
{
getTokenCacheKey
(
token
.
Key
),
getTokenCacheFenceKey
(
token
.
Key
),
},
token
.
Id
,
token
.
UserId
,
token
.
Status
,
token
.
Name
,
token
.
CreatedTime
,
token
.
AccessedTime
,
token
.
ExpiredTime
,
strconv
.
FormatBool
(
token
.
UnlimitedQuota
),
strconv
.
FormatBool
(
token
.
ModelLimitsEnabled
),
token
.
ModelLimits
,
allowIps
,
token
.
Group
,
strconv
.
FormatBool
(
token
.
CrossGroupRetry
),
token
.
AutoGroups
,
token
.
RemainQuota
,
token
.
UsedQuota
,
tokenCacheTTLSeconds
(),
)
.
Int
()
}
// cacheGetTokenByKey 从缓存读取 token;不完整的哈希(如仅有配额字段)会被拒绝。
func
cacheGetTokenByKey
(
key
string
)
(
*
Token
,
error
)
{
func
cacheGetTokenByKey
(
key
string
)
(
*
Token
,
error
)
{
hmacKey
:=
common
.
GenerateHMAC
(
key
)
if
!
common
.
RedisEnabled
{
if
!
common
.
RedisEnabled
{
return
nil
,
fmt
.
Errorf
(
"redis is not enabled"
)
return
nil
,
fmt
.
Errorf
(
"redis is not enabled"
)
}
}
var
token
Token
var
token
Token
err
:=
common
.
RedisHGetObj
(
fmt
.
Sprintf
(
"token:%s"
,
hmacKey
),
&
token
)
if
err
:=
common
.
RedisHGetObj
(
getTokenCacheKey
(
key
),
&
token
);
err
!=
nil
{
if
err
!=
nil
{
return
nil
,
err
return
nil
,
err
}
}
if
token
.
Id
<=
0
{
return
nil
,
fmt
.
Errorf
(
"token cache is incomplete"
)
}
token
.
Key
=
key
token
.
Key
=
key
return
&
token
,
nil
return
&
token
,
nil
}
}
model/user.go
View file @
ccd535ef
...
@@ -547,7 +547,7 @@ func inviteUser(inviterId int) error {
...
@@ -547,7 +547,7 @@ func inviteUser(inviterId int) error {
func
(
user
*
User
)
TransferAffQuotaToQuota
(
quota
int
)
error
{
func
(
user
*
User
)
TransferAffQuotaToQuota
(
quota
int
)
error
{
// 检查quota是否小于最小额度
// 检查quota是否小于最小额度
if
float64
(
quota
)
<
common
.
QuotaPerUnit
{
if
float64
(
quota
)
<
common
.
QuotaPerUnit
{
return
fmt
.
Errorf
(
"转移额度最小为%s!"
,
logger
.
LogQuota
(
in
t
(
common
.
QuotaPerUnit
)))
return
fmt
.
Errorf
(
"转移额度最小为%s!"
,
logger
.
LogQuota
(
common
.
QuotaFromFloa
t
(
common
.
QuotaPerUnit
)))
}
}
// 开始数据库事务
// 开始数据库事务
...
@@ -1180,24 +1180,9 @@ func ValidateAccessToken(token string) (*User, error) {
...
@@ -1180,24 +1180,9 @@ func ValidateAccessToken(token string) (*User, error) {
// GetUserQuota gets quota from Redis first, falls back to DB if needed
// GetUserQuota gets quota from Redis first, falls back to DB if needed
func
GetUserQuota
(
id
int
,
fromDB
bool
)
(
quota
int
,
err
error
)
{
func
GetUserQuota
(
id
int
,
fromDB
bool
)
(
quota
int
,
err
error
)
{
defer
func
()
{
// Update Redis cache asynchronously on successful DB read
if
shouldUpdateRedis
(
fromDB
,
err
)
{
gopool
.
Go
(
func
()
{
if
err
:=
updateUserQuotaCache
(
id
,
quota
);
err
!=
nil
{
common
.
SysLog
(
"failed to update user quota cache: "
+
err
.
Error
())
}
})
}
}()
if
!
fromDB
&&
common
.
RedisEnabled
{
if
!
fromDB
&&
common
.
RedisEnabled
{
quota
,
err
:=
getUserQuotaCache
(
id
)
return
getUserQuotaCache
(
id
)
if
err
==
nil
{
return
quota
,
nil
}
// Don't return error - fall through to DB
}
}
fromDB
=
true
err
=
DB
.
Model
(
&
User
{})
.
Where
(
"id = ?"
,
id
)
.
Select
(
"quota"
)
.
Find
(
&
quota
)
.
Error
err
=
DB
.
Model
(
&
User
{})
.
Where
(
"id = ?"
,
id
)
.
Select
(
"quota"
)
.
Find
(
&
quota
)
.
Error
if
err
!=
nil
{
if
err
!=
nil
{
return
0
,
err
return
0
,
err
...
@@ -1403,24 +1388,6 @@ func updateUserQuotaUsedQuotaAndRequestCount(id int, quota int, usedQuota int, r
...
@@ -1403,24 +1388,6 @@ func updateUserQuotaUsedQuotaAndRequestCount(id int, quota int, usedQuota int, r
}
}
}
}
func
updateUserUsedQuota
(
id
int
,
quota
int
)
{
err
:=
DB
.
Model
(
&
User
{})
.
Where
(
"id = ?"
,
id
)
.
Updates
(
map
[
string
]
interface
{}{
"used_quota"
:
gorm
.
Expr
(
"used_quota + ?"
,
quota
),
},
)
.
Error
if
err
!=
nil
{
common
.
SysLog
(
"failed to update user used quota: "
+
err
.
Error
())
}
}
func
updateUserRequestCount
(
id
int
,
count
int
)
{
err
:=
DB
.
Model
(
&
User
{})
.
Where
(
"id = ?"
,
id
)
.
Update
(
"request_count"
,
gorm
.
Expr
(
"request_count + ?"
,
count
))
.
Error
if
err
!=
nil
{
common
.
SysLog
(
"failed to update user request count: "
+
err
.
Error
())
}
}
// GetUsernameById gets username from Redis first, falls back to DB if needed
// GetUsernameById gets username from Redis first, falls back to DB if needed
func
GetUsernameById
(
id
int
,
fromDB
bool
)
(
username
string
,
err
error
)
{
func
GetUsernameById
(
id
int
,
fromDB
bool
)
(
username
string
,
err
error
)
{
defer
func
()
{
defer
func
()
{
...
...
model/user_cache.go
View file @
ccd535ef
...
@@ -67,12 +67,6 @@ func invalidateUserCache(userId int) error {
...
@@ -67,12 +67,6 @@ func invalidateUserCache(userId int) error {
return
common
.
RedisDelKey
(
getUserCacheKey
(
userId
))
return
common
.
RedisDelKey
(
getUserCacheKey
(
userId
))
}
}
// InvalidateUserCache is the exported version of invalidateUserCache.
// 供 controller 等上层包在用户状态变更(如禁用、删除、角色变更)后主动清理缓存。
func
InvalidateUserCache
(
userId
int
)
error
{
return
invalidateUserCache
(
userId
)
}
func
populateUserCache
(
user
User
)
error
{
func
populateUserCache
(
user
User
)
error
{
if
!
common
.
RedisEnabled
{
if
!
common
.
RedisEnabled
{
return
nil
return
nil
...
@@ -187,14 +181,6 @@ func getUserQuotaCache(userId int) (int, error) {
...
@@ -187,14 +181,6 @@ func getUserQuotaCache(userId int) (int, error) {
return
cache
.
Quota
,
nil
return
cache
.
Quota
,
nil
}
}
func
getUserStatusCache
(
userId
int
)
(
int
,
error
)
{
cache
,
err
:=
GetUserCache
(
userId
)
if
err
!=
nil
{
return
0
,
err
}
return
cache
.
Status
,
nil
}
func
getUserNameCache
(
userId
int
)
(
string
,
error
)
{
func
getUserNameCache
(
userId
int
)
(
string
,
error
)
{
cache
,
err
:=
GetUserCache
(
userId
)
cache
,
err
:=
GetUserCache
(
userId
)
if
err
!=
nil
{
if
err
!=
nil
{
...
@@ -211,22 +197,6 @@ func getUserSettingCache(userId int) (dto.UserSetting, error) {
...
@@ -211,22 +197,6 @@ func getUserSettingCache(userId int) (dto.UserSetting, error) {
return
cache
.
GetSetting
(),
nil
return
cache
.
GetSetting
(),
nil
}
}
// New functions for individual field updates
func
updateUserStatusCache
(
userId
int
,
status
bool
)
error
{
statusInt
:=
common
.
UserStatusEnabled
if
!
status
{
statusInt
=
common
.
UserStatusDisabled
}
return
updateUserCacheField
(
userId
,
"Status"
,
statusInt
)
}
func
updateUserQuotaCache
(
userId
int
,
quota
int
)
error
{
if
!
common
.
RedisEnabled
{
return
nil
}
return
common
.
RedisHSetField
(
getUserCacheKey
(
userId
),
"Quota"
,
fmt
.
Sprintf
(
"%d"
,
quota
))
}
// RefreshUserGroupCache writes the database-authoritative group into an
// RefreshUserGroupCache writes the database-authoritative group into an
// existing user hash without changing the user's authentication version.
// existing user hash without changing the user's authentication version.
func
RefreshUserGroupCache
(
userId
int
)
error
{
func
RefreshUserGroupCache
(
userId
int
)
error
{
...
...
service/billing_session.go
View file @
ccd535ef
package
service
package
service
import
(
import
(
"errors"
"fmt"
"fmt"
"net/http"
"net/http"
"strings"
"strings"
...
@@ -214,6 +215,16 @@ func (s *BillingSession) preConsume(c *gin.Context, quota int) *types.NewAPIErro
...
@@ -214,6 +215,16 @@ func (s *BillingSession) preConsume(c *gin.Context, quota int) *types.NewAPIErro
s
.
tokenConsumed
=
0
s
.
tokenConsumed
=
0
}
}
// TODO: model 层应定义哨兵错误(如 ErrNoActiveSubscription),用 errors.Is 替代字符串匹配
// TODO: model 层应定义哨兵错误(如 ErrNoActiveSubscription),用 errors.Is 替代字符串匹配
if
errors
.
Is
(
err
,
ErrInsufficientWalletQuota
)
{
userQuota
,
quotaErr
:=
model
.
GetUserQuota
(
s
.
relayInfo
.
UserId
,
false
)
if
quotaErr
!=
nil
{
userQuota
=
0
}
return
types
.
NewErrorWithStatusCode
(
fmt
.
Errorf
(
"用户额度不足, 剩余额度: %s"
,
logger
.
FormatQuota
(
userQuota
)),
types
.
ErrorCodeInsufficientUserQuota
,
http
.
StatusForbidden
,
types
.
ErrOptionWithSkipRetry
(),
types
.
ErrOptionWithNoRecordErrorLog
())
}
errMsg
:=
err
.
Error
()
errMsg
:=
err
.
Error
()
if
strings
.
Contains
(
errMsg
,
"no active subscription"
)
||
strings
.
Contains
(
errMsg
,
"subscription quota insufficient"
)
{
if
strings
.
Contains
(
errMsg
,
"no active subscription"
)
||
strings
.
Contains
(
errMsg
,
"subscription quota insufficient"
)
{
return
types
.
NewErrorWithStatusCode
(
fmt
.
Errorf
(
"订阅额度不足或未配置订阅: %s"
,
errMsg
),
types
.
ErrorCodeInsufficientUserQuota
,
http
.
StatusForbidden
,
types
.
ErrOptionWithSkipRetry
(),
types
.
ErrOptionWithNoRecordErrorLog
())
return
types
.
NewErrorWithStatusCode
(
fmt
.
Errorf
(
"订阅额度不足或未配置订阅: %s"
,
errMsg
),
types
.
ErrorCodeInsufficientUserQuota
,
http
.
StatusForbidden
,
types
.
ErrOptionWithSkipRetry
(),
types
.
ErrOptionWithNoRecordErrorLog
())
...
...
service/funding_source.go
View file @
ccd535ef
package
service
package
service
import
(
import
(
"errors"
"time"
"time"
"github.com/QuantumNous/new-api/model"
"github.com/QuantumNous/new-api/model"
...
@@ -26,6 +27,11 @@ type FundingSource interface {
...
@@ -26,6 +27,11 @@ type FundingSource interface {
// WalletFunding — 钱包资金来源实现
// WalletFunding — 钱包资金来源实现
// ---------------------------------------------------------------------------
// ---------------------------------------------------------------------------
// ErrInsufficientWalletQuota 钱包原子预扣失败(余额不足),未发生任何扣减。
// BillingSession 据此映射为 ErrorCodeInsufficientUserQuota,
// 使 wallet_first 等计费偏好可以回退到订阅。
var
ErrInsufficientWalletQuota
=
errors
.
New
(
"wallet quota insufficient"
)
type
WalletFunding
struct
{
type
WalletFunding
struct
{
userId
int
userId
int
consumed
int
// 实际预扣的用户额度
consumed
int
// 实际预扣的用户额度
...
@@ -37,9 +43,13 @@ func (w *WalletFunding) PreConsume(amount int) error {
...
@@ -37,9 +43,13 @@ func (w *WalletFunding) PreConsume(amount int) error {
if
amount
<=
0
{
if
amount
<=
0
{
return
nil
return
nil
}
}
if
err
:=
model
.
DecreaseUserQuota
(
w
.
userId
,
amount
,
false
);
err
!=
nil
{
reserved
,
err
:=
model
.
TryReserveUserQuota
(
w
.
userId
,
amount
)
if
err
!=
nil
{
return
err
return
err
}
}
if
!
reserved
{
return
ErrInsufficientWalletQuota
}
w
.
consumed
=
amount
w
.
consumed
=
amount
return
nil
return
nil
}
}
...
...
service/quota.go
View file @
ccd535ef
...
@@ -391,19 +391,17 @@ func PreConsumeTokenQuota(relayInfo *relaycommon.RelayInfo, quota int) error {
...
@@ -391,19 +391,17 @@ func PreConsumeTokenQuota(relayInfo *relaycommon.RelayInfo, quota int) error {
if
relayInfo
.
IsPlayground
{
if
relayInfo
.
IsPlayground
{
return
nil
return
nil
}
}
//if relayInfo.TokenUnlimited {
// 原子预扣:检查与扣减在同一操作中完成,并发请求不可能同时通过检查后超扣。
// return nil
reserved
,
err
:=
model
.
TryReserveTokenQuota
(
relayInfo
.
TokenId
,
relayInfo
.
TokenKey
,
quota
,
relayInfo
.
TokenUnlimited
)
//}
token
,
err
:=
model
.
GetTokenByKey
(
relayInfo
.
TokenKey
,
false
)
if
err
!=
nil
{
if
err
!=
nil
{
return
err
return
err
}
}
if
!
re
layInfo
.
TokenUnlimited
&&
token
.
RemainQuota
<
quota
{
if
!
re
served
{
re
turn
fmt
.
Errorf
(
"token quota is not enough, token remain quota: %s, need quota: %s"
,
logger
.
FormatQuota
(
token
.
RemainQuota
),
logger
.
FormatQuota
(
quota
))
re
mainQuota
:=
0
}
if
token
,
tokenErr
:=
model
.
GetTokenByKey
(
relayInfo
.
TokenKey
,
false
);
tokenErr
==
nil
&&
token
!=
nil
{
err
=
model
.
DecreaseTokenQuota
(
relayInfo
.
TokenId
,
relayInfo
.
TokenKey
,
quota
)
remainQuota
=
token
.
RemainQuota
if
err
!=
nil
{
}
return
err
return
fmt
.
Errorf
(
"token quota is not enough, token remain quota: %s, need quota: %s"
,
logger
.
FormatQuota
(
remainQuota
),
logger
.
FormatQuota
(
quota
))
}
}
return
nil
return
nil
}
}
...
...
Write
Preview
Markdown
is supported
0%
Try again
or
attach a new file
Attach a file
Cancel
You are about to add
0
people
to the discussion. Proceed with caution.
Finish editing this message first!
Cancel
Please
register
or
sign in
to comment