Merge pull request #5005 from Wei-Shaw/refactor/scoped-column-updates

refactor(repository): scope user/api-key updates to explicitly declared columns
This commit is contained in:
Wesley Liddick
2026-07-28 19:43:50 +08:00
committed by GitHub
41 changed files with 1159 additions and 177 deletions
@@ -3015,7 +3015,7 @@ func (r *oauthPendingFlowUserRepo) GetFirstAdmin(context.Context) (*service.User
panic("unexpected GetFirstAdmin call")
}
func (r *oauthPendingFlowUserRepo) Update(ctx context.Context, user *service.User) error {
func (r *oauthPendingFlowUserRepo) Update(ctx context.Context, user *service.User, fields service.UserUpdateFields) error {
entity, err := r.client.User.UpdateOneID(user.ID).
SetEmail(user.Email).
SetUsername(user.Username).
@@ -3168,6 +3168,14 @@ func (r *oauthPendingFlowUserRepo) DeductBalance(context.Context, int64, float64
panic("unexpected DeductBalance call")
}
func (r *oauthPendingFlowUserRepo) AdjustBalance(ctx context.Context, id int64, delta float64) (service.BalanceChange, error) {
panic("unexpected AdjustBalance call")
}
func (r *oauthPendingFlowUserRepo) SetBalance(ctx context.Context, id int64, value float64) (service.BalanceChange, error) {
panic("unexpected SetBalance call")
}
func (r *oauthPendingFlowUserRepo) UpdateConcurrency(context.Context, int64, int) error {
panic("unexpected UpdateConcurrency call")
}
@@ -47,7 +47,11 @@ func TestAuthHandlerRevokeAllSessionsInvalidatesAccessTokens(t *testing.T) {
require.Equal(t, http.StatusOK, recorder.Code)
require.Equal(t, []int64{29}, refreshTokenCache.revokedUserIDs)
require.Equal(t, int64(8), repo.user.TokenVersion)
// users 表没有 token_version 列(见 resolvedTokenVersion:JWT 里的值由
// email+password_hash 指纹推导),所以自增 TokenVersion 只停留在内存里。
// 此前紧跟其后的整行 Update 不写任何有效数据,却会用旧快照覆盖并发写入的列,
// 已移除。会话撤销由上面的 refresh session 清理承担。
require.Equal(t, int64(7), repo.user.TokenVersion)
var resp struct {
Code int `json:"code"`
+14 -2
View File
@@ -41,7 +41,7 @@ func (s *userHandlerRepoStub) GetFirstAdmin(context.Context) (*service.User, err
cloned := *s.user
return &cloned, nil
}
func (s *userHandlerRepoStub) Update(_ context.Context, user *service.User) error {
func (s *userHandlerRepoStub) Update(_ context.Context, user *service.User, _ service.UserUpdateFields) error {
cloned := *user
s.user = &cloned
return nil
@@ -93,6 +93,14 @@ func (s *userHandlerRepoStub) UpdateConcurrency(context.Context, int64, int) err
func (s *userHandlerRepoStub) BatchSetConcurrency(context.Context, []int64, int) (int, error) {
return 0, nil
}
func (s *userHandlerRepoStub) AdjustBalance(ctx context.Context, id int64, delta float64) (service.BalanceChange, error) {
panic("unexpected AdjustBalance call")
}
func (s *userHandlerRepoStub) SetBalance(ctx context.Context, id int64, value float64) (service.BalanceChange, error) {
panic("unexpected SetBalance call")
}
func (s *userHandlerRepoStub) BatchAddConcurrency(context.Context, []int64, int) (int, error) {
return 0, nil
}
@@ -656,7 +664,11 @@ func TestUserHandlerUnbindIdentityRevokesAllUserSessionsWhenAuthServiceConfigure
require.Equal(t, http.StatusOK, recorder.Code)
require.Equal(t, []int64{23}, refreshTokenCache.revokedUserIDs)
require.Equal(t, int64(5), repo.user.TokenVersion)
// 撤销依赖的是 refresh session 清理,而不是 token_version:users 表没有这一列
// (见 resolvedTokenVersion,实际值由 email+password_hash 指纹推导),
// 所以此前"自增 TokenVersion 再整行写回"不持久化任何东西,
// 却会用旧快照覆盖并发写入的列。这里断言用户行未被改写。
require.Equal(t, int64(4), repo.user.TokenVersion)
}
func TestUserHandlerUnbindIdentityDoesNotRevokeSessionsWhenNothingWasUnbound(t *testing.T) {
+70 -45
View File
@@ -222,7 +222,12 @@ func (r *apiKeyRepository) GetByKeyForAuth(ctx context.Context, key string) (*se
return apiKeyEntityToService(m), nil
}
func (r *apiKeyRepository) Update(ctx context.Context, key *service.APIKey) error {
func (r *apiKeyRepository) Update(ctx context.Context, key *service.APIKey, fields service.APIKeyUpdateFields) error {
// 空掩码代表调用方不改任何列,直接返回,避免产生一次无意义的整行写。
if fields.IsEmpty() {
return nil
}
// 使用原子操作:将软删除检查与更新合并到同一语句,避免竞态条件。
// 之前的实现先检查 Exist 再 UpdateOneID,若在两步之间发生软删除,
// 则会更新已删除的记录。
@@ -232,57 +237,77 @@ func (r *apiKeyRepository) Update(ctx context.Context, key *service.APIKey) erro
now := time.Now()
builder := client.APIKey.Update().
Where(apikey.IDEQ(key.ID), apikey.DeletedAtIsNil()).
SetName(key.Name).
SetStatus(key.Status).
SetQuota(key.Quota).
SetQuotaUsed(key.QuotaUsed).
SetRateLimit5h(key.RateLimit5h).
SetRateLimit1d(key.RateLimit1d).
SetRateLimit7d(key.RateLimit7d).
SetUsage5h(key.Usage5h).
SetUsage1d(key.Usage1d).
SetUsage7d(key.Usage7d).
SetUpdatedAt(now)
if key.GroupID != nil {
builder.SetGroupID(*key.GroupID)
} else {
builder.ClearGroupID()
if fields.Name {
builder.SetName(key.Name)
}
if fields.Status {
builder.SetStatus(key.Status)
}
if fields.Quota {
builder.SetQuota(key.Quota)
}
if fields.QuotaUsed {
builder.SetQuotaUsed(key.QuotaUsed)
}
if fields.RateLimits {
builder.
SetRateLimit5h(key.RateLimit5h).
SetRateLimit1d(key.RateLimit1d).
SetRateLimit7d(key.RateLimit7d)
}
if fields.RateLimitUsage {
builder.
SetUsage5h(key.Usage5h).
SetUsage1d(key.Usage1d).
SetUsage7d(key.Usage7d)
// Rate limit window start times
if key.Window5hStart != nil {
builder.SetWindow5hStart(*key.Window5hStart)
} else {
builder.ClearWindow5hStart()
}
if key.Window1dStart != nil {
builder.SetWindow1dStart(*key.Window1dStart)
} else {
builder.ClearWindow1dStart()
}
if key.Window7dStart != nil {
builder.SetWindow7dStart(*key.Window7dStart)
} else {
builder.ClearWindow7dStart()
}
}
if fields.GroupID {
if key.GroupID != nil {
builder.SetGroupID(*key.GroupID)
} else {
builder.ClearGroupID()
}
}
// Expiration time
if key.ExpiresAt != nil {
builder.SetExpiresAt(*key.ExpiresAt)
} else {
builder.ClearExpiresAt()
}
// Rate limit window start times
if key.Window5hStart != nil {
builder.SetWindow5hStart(*key.Window5hStart)
} else {
builder.ClearWindow5hStart()
}
if key.Window1dStart != nil {
builder.SetWindow1dStart(*key.Window1dStart)
} else {
builder.ClearWindow1dStart()
}
if key.Window7dStart != nil {
builder.SetWindow7dStart(*key.Window7dStart)
} else {
builder.ClearWindow7dStart()
if fields.ExpiresAt {
if key.ExpiresAt != nil {
builder.SetExpiresAt(*key.ExpiresAt)
} else {
builder.ClearExpiresAt()
}
}
// IP 限制字段
if len(key.IPWhitelist) > 0 {
builder.SetIPWhitelist(key.IPWhitelist)
} else {
builder.ClearIPWhitelist()
}
if len(key.IPBlacklist) > 0 {
builder.SetIPBlacklist(key.IPBlacklist)
} else {
builder.ClearIPBlacklist()
if fields.IPRules {
if len(key.IPWhitelist) > 0 {
builder.SetIPWhitelist(key.IPWhitelist)
} else {
builder.ClearIPWhitelist()
}
if len(key.IPBlacklist) > 0 {
builder.SetIPBlacklist(key.IPBlacklist)
} else {
builder.ClearIPBlacklist()
}
}
affected, err := builder.Save(ctx)
@@ -139,7 +139,7 @@ func (s *APIKeyRepoSuite) TestUpdate() {
key.Name = "Renamed"
key.Status = service.StatusDisabled
err := s.repo.Update(s.ctx, key)
err := s.repo.Update(s.ctx, key, service.APIKeyUpdateFields{Name: true, Status: true})
s.Require().NoError(err, "Update")
got, err := s.repo.GetByID(s.ctx, key.ID)
@@ -163,7 +163,7 @@ func (s *APIKeyRepoSuite) TestUpdate_ClearGroupID() {
s.Require().NoError(s.repo.Create(s.ctx, key))
key.GroupID = nil
err := s.repo.Update(s.ctx, key)
err := s.repo.Update(s.ctx, key, service.APIKeyUpdateFields{GroupID: true})
s.Require().NoError(err, "Update")
got, err := s.repo.GetByID(s.ctx, key.ID)
@@ -368,7 +368,7 @@ func (s *APIKeyRepoSuite) TestCRUD_Search_ClearGroupID() {
key.Name = "Renamed"
key.Status = service.StatusDisabled
key.GroupID = nil
s.Require().NoError(s.repo.Update(s.ctx, key), "Update")
s.Require().NoError(s.repo.Update(s.ctx, key, service.APIKeyUpdateFields{Name: true, Status: true, GroupID: true}), "Update")
got2, err := s.repo.GetByID(s.ctx, key.ID)
s.Require().NoError(err, "GetByID")
@@ -486,7 +486,7 @@ func (s *APIKeyRepoSuite) TestIncrementQuotaUsedAndGetState() {
key := s.mustCreateApiKey(user.ID, "sk-quota-state", "QuotaState", nil)
key.Quota = 3
key.QuotaUsed = 1
s.Require().NoError(s.repo.Update(s.ctx, key), "Update quota")
s.Require().NoError(s.repo.Update(s.ctx, key, service.APIKeyUpdateFields{Quota: true, QuotaUsed: true}), "Update quota")
state, err := s.repo.IncrementQuotaUsedAndGetState(s.ctx, key.ID, 2.5)
s.Require().NoError(err, "IncrementQuotaUsedAndGetState")
@@ -0,0 +1,111 @@
//go:build integration
package repository
import (
"github.com/Wei-Shaw/sub2api/internal/service"
)
// api_keys 上的用量列由计费热路径原子递增(IncrementQuotaUsed /
// IncrementRateLimitUsage)。编辑 Key 时若整行回写,
// 并发累计的配额与限流计数就会被旧快照覆盖。
func (s *APIKeyRepoSuite) TestUpdate_DoesNotRevertConcurrentQuotaUsage() {
user := s.mustCreateUser("apikey-lost-update-quota@example.com")
key := &service.APIKey{
UserID: user.ID,
Key: "sk-lost-update-quota",
Name: "before",
Status: service.StatusActive,
Quota: 100,
}
s.Require().NoError(s.repo.Create(s.ctx, key), "Create")
stale, err := s.repo.GetByID(s.ctx, key.ID)
s.Require().NoError(err, "GetByID")
s.Require().Zero(stale.QuotaUsed)
newUsed, err := s.repo.IncrementQuotaUsed(s.ctx, key.ID, 30)
s.Require().NoError(err, "IncrementQuotaUsed")
s.Require().InDelta(30, newUsed, 1e-9)
stale.Name = "after"
s.Require().NoError(
s.repo.Update(s.ctx, stale, service.APIKeyUpdateFields{Name: true}),
"Update",
)
got, err := s.repo.GetByID(s.ctx, key.ID)
s.Require().NoError(err, "GetByID after update")
s.Require().Equal("after", got.Name, "declared column must still be written")
s.Require().InDelta(30, got.QuotaUsed, 1e-9, "quota_used must not be reverted by a stale key edit")
}
func (s *APIKeyRepoSuite) TestUpdate_DoesNotRevertConcurrentRateLimitUsage() {
user := s.mustCreateUser("apikey-lost-update-ratelimit@example.com")
key := &service.APIKey{
UserID: user.ID,
Key: "sk-lost-update-ratelimit",
Name: "before",
Status: service.StatusActive,
RateLimit5h: 100,
}
s.Require().NoError(s.repo.Create(s.ctx, key), "Create")
stale, err := s.repo.GetByID(s.ctx, key.ID)
s.Require().NoError(err, "GetByID")
s.Require().Zero(stale.Usage5h)
s.Require().NoError(s.repo.IncrementRateLimitUsage(s.ctx, key.ID, 42), "IncrementRateLimitUsage")
stale.Name = "after"
s.Require().NoError(
s.repo.Update(s.ctx, stale, service.APIKeyUpdateFields{Name: true}),
"Update",
)
got, err := s.repo.GetByID(s.ctx, key.ID)
s.Require().NoError(err, "GetByID after update")
s.Require().InDelta(42, got.Usage5h, 1e-9, "usage_5h must not be reverted by a stale key edit")
s.Require().InDelta(42, got.Usage1d, 1e-9, "usage_1d must not be reverted by a stale key edit")
s.Require().InDelta(42, got.Usage7d, 1e-9, "usage_7d must not be reverted by a stale key edit")
}
// 显式重置仍然必须生效,避免收窄写入列时把功能改坏。
func (s *APIKeyRepoSuite) TestUpdate_StillResetsUsageWhenDeclared() {
user := s.mustCreateUser("apikey-reset-usage@example.com")
key := &service.APIKey{
UserID: user.ID,
Key: "sk-reset-usage",
Name: "reset",
Status: service.StatusActive,
Quota: 100,
}
s.Require().NoError(s.repo.Create(s.ctx, key), "Create")
_, err := s.repo.IncrementQuotaUsed(s.ctx, key.ID, 30)
s.Require().NoError(err, "IncrementQuotaUsed")
s.Require().NoError(s.repo.IncrementRateLimitUsage(s.ctx, key.ID, 42), "IncrementRateLimitUsage")
current, err := s.repo.GetByID(s.ctx, key.ID)
s.Require().NoError(err, "GetByID")
current.QuotaUsed = 0
current.Usage5h = 0
current.Usage1d = 0
current.Usage7d = 0
current.Window5hStart = nil
current.Window1dStart = nil
current.Window7dStart = nil
s.Require().NoError(
s.repo.Update(s.ctx, current, service.APIKeyUpdateFields{QuotaUsed: true, RateLimitUsage: true}),
"Update",
)
got, err := s.repo.GetByID(s.ctx, key.ID)
s.Require().NoError(err, "GetByID after reset")
s.Require().Zero(got.QuotaUsed, "explicit quota reset must still apply")
s.Require().Zero(got.Usage5h, "explicit rate limit reset must still apply")
s.Require().Zero(got.Usage1d)
s.Require().Zero(got.Usage7d)
s.Require().Nil(got.Window5hStart)
}
@@ -77,8 +77,8 @@ func TestAuthCacheInvalidationTriggers_CoverSecurityMutationsOnly(t *testing.T)
userRepo := NewUserRepository(integrationEntClient, integrationDB)
loadedUser, err := userRepo.GetByID(ctx, user.ID)
require.NoError(t, err)
loadedUser.Balance += 10
require.NoError(t, userRepo.Update(ctx, loadedUser))
_, err = userRepo.AdjustBalance(ctx, loadedUser.ID, 10)
require.NoError(t, err)
require.Zero(t, count(), "balance update with unchanged allowed groups must not enqueue")
_, err = integrationDB.ExecContext(ctx, "UPDATE users SET status = 'disabled' WHERE id = $1", user.ID)
@@ -87,13 +87,18 @@ func (r *promoCodeRepository) GetByCodeForUpdate(ctx context.Context, code strin
return promoCodeEntityToService(m), nil
}
// Update 写入管理员可编辑的字段。
//
// 这里刻意不写 used_count:它由兑换路径的 IncrementUsedCount 原子递增,
// 而 used_count >= max_uses 正是"优惠码用完了"的判定依据。若管理员编辑
// (改有效期、改额度……)时按快照把 used_count 回写,并发的兑换计数就会被抹掉,
// 兑换次数统计随之失真。PromoService.Update 也从不修改该字段。
func (r *promoCodeRepository) Update(ctx context.Context, code *service.PromoCode) error {
client := clientFromContext(ctx, r.client)
builder := client.PromoCode.UpdateOneID(code.ID).
SetCode(code.Code).
SetBonusAmount(code.BonusAmount).
SetMaxUses(code.MaxUses).
SetUsedCount(code.UsedCount).
SetStatus(code.Status).
SetNotes(code.Notes)
@@ -482,7 +482,8 @@ func (s *UserProfileIdentityRepoSuite) TestWithUserProfileIdentityTx_AllowsAvata
if err != nil {
return err
}
return s.repo.Update(txCtx, model)
// 只改头像时用户行没有任何列需要写,掩码为空——与 UserService.updateProfile 一致。
return s.repo.Update(txCtx, model, service.UserUpdateFields{})
})
s.Require().NoError(err)
+167 -36
View File
@@ -205,10 +205,14 @@ func (r *userRepository) GetByEmail(ctx context.Context, email string) (*service
return out, nil
}
func (r *userRepository) Update(ctx context.Context, userIn *service.User) error {
func (r *userRepository) Update(ctx context.Context, userIn *service.User, fields service.UserUpdateFields) error {
if userIn == nil {
return nil
}
// 空掩码代表调用方不改任何列,直接返回,避免产生一次无意义的整行写。
if fields.IsEmpty() {
return nil
}
// 使用 ent 事务包裹用户更新与 allowed_groups 同步,避免跨层事务不一致。
tx, err := r.client.Tx(ctx)
@@ -231,19 +235,23 @@ func (r *userRepository) Update(ctx context.Context, userIn *service.User) error
}
}
releaseEmailLock, err := lockRepositoryScopedKeys(
txCtx,
txClient,
txAwareSQLExecutor(txCtx, r.sql, r.client),
normalizedEmailUniquenessLockKey(userIn.Email),
)
if err != nil {
return err
}
defer releaseEmailLock()
// 邮箱唯一性锁与查重只在本次确实要改邮箱时才做:不改邮箱的更新既不需要
// 串行化,也不该因为快照里的旧邮箱已被他人占用而报 ErrEmailExists。
if fields.Email {
releaseEmailLock, err := lockRepositoryScopedKeys(
txCtx,
txClient,
txAwareSQLExecutor(txCtx, r.sql, r.client),
normalizedEmailUniquenessLockKey(userIn.Email),
)
if err != nil {
return err
}
defer releaseEmailLock()
if err := ensureNormalizedEmailAvailableWithClient(txCtx, txClient, userIn.ID, userIn.Email); err != nil {
return err
if err := ensureNormalizedEmailAvailableWithClient(txCtx, txClient, userIn.ID, userIn.Email); err != nil {
return err
}
}
existing, err := clientFromContext(txCtx, txClient).User.Get(txCtx, userIn.ID)
@@ -252,41 +260,64 @@ func (r *userRepository) Update(ctx context.Context, userIn *service.User) error
}
oldEmail := existing.Email
updateOp := txClient.User.UpdateOneID(userIn.ID).
SetEmail(userIn.Email).
SetUsername(userIn.Username).
SetNotes(userIn.Notes).
SetPasswordHash(userIn.PasswordHash).
SetRole(userIn.Role).
SetBalance(userIn.Balance).
SetConcurrency(userIn.Concurrency).
SetStatus(userIn.Status).
SetBalanceNotifyEnabled(userIn.BalanceNotifyEnabled).
SetBalanceNotifyThresholdType(userIn.BalanceNotifyThresholdType).
SetNillableBalanceNotifyThreshold(userIn.BalanceNotifyThreshold).
SetBalanceNotifyExtraEmails(marshalExtraEmails(userIn.BalanceNotifyExtraEmails)).
SetTotalRecharged(userIn.TotalRecharged).
SetRpmLimit(userIn.RPMLimit)
if userIn.SignupSource != "" {
updateOp := txClient.User.UpdateOneID(userIn.ID)
if fields.Email {
updateOp = updateOp.SetEmail(userIn.Email)
}
if fields.Username {
updateOp = updateOp.SetUsername(userIn.Username)
}
if fields.Notes {
updateOp = updateOp.SetNotes(userIn.Notes)
}
if fields.PasswordHash {
updateOp = updateOp.SetPasswordHash(userIn.PasswordHash)
}
if fields.Role {
updateOp = updateOp.SetRole(userIn.Role)
}
if fields.Concurrency {
updateOp = updateOp.SetConcurrency(userIn.Concurrency)
}
if fields.RPMLimit {
updateOp = updateOp.SetRpmLimit(userIn.RPMLimit)
}
if fields.Status {
updateOp = updateOp.SetStatus(userIn.Status)
}
if fields.BalanceNotifySettings {
updateOp = updateOp.
SetBalanceNotifyEnabled(userIn.BalanceNotifyEnabled).
SetBalanceNotifyThresholdType(userIn.BalanceNotifyThresholdType).
SetNillableBalanceNotifyThreshold(userIn.BalanceNotifyThreshold)
if userIn.BalanceNotifyThreshold == nil {
updateOp = updateOp.ClearBalanceNotifyThreshold()
}
}
if fields.BalanceNotifyExtraEmails {
updateOp = updateOp.SetBalanceNotifyExtraEmails(marshalExtraEmails(userIn.BalanceNotifyExtraEmails))
}
if fields.SignupSource && userIn.SignupSource != "" {
updateOp = updateOp.SetSignupSource(userIn.SignupSource)
}
if userIn.LastLoginAt != nil {
if fields.LastLoginAt && userIn.LastLoginAt != nil {
updateOp = updateOp.SetLastLoginAt(*userIn.LastLoginAt)
}
if userIn.LastActiveAt != nil {
if fields.LastActiveAt && userIn.LastActiveAt != nil {
updateOp = updateOp.SetLastActiveAt(*userIn.LastActiveAt)
}
if userIn.BalanceNotifyThreshold == nil {
updateOp = updateOp.ClearBalanceNotifyThreshold()
}
updated, err := updateOp.Save(txCtx)
if err != nil {
return translatePersistenceError(err, service.ErrUserNotFound, service.ErrEmailExists)
}
if err := r.syncUserAllowedGroupsWithClient(txCtx, txClient, updated.ID, userIn.AllowedGroups); err != nil {
return err
if fields.AllowedGroups {
if err := r.syncUserAllowedGroupsWithClient(txCtx, txClient, updated.ID, userIn.AllowedGroups); err != nil {
return err
}
}
// 始终以库中的邮箱为准补齐 email 身份:未改邮箱时 updated.Email == oldEmail,
// 这里退化为幂等的身份补写,与改邮箱前的行为一致。
if err := replaceEmailAuthIdentityWithClient(txCtx, txClient, updated.ID, oldEmail, updated.Email, "user_repo_update"); err != nil {
return err
}
@@ -828,6 +859,106 @@ func (r *userRepository) DeductBalance(ctx context.Context, id int64, amount flo
return nil
}
// AdjustBalance 原子地把 delta 累加到余额上,结果为负时整条语句不生效。
// 相比"读余额 → 算新值 → 整行写回",这里把读与写压进同一条 UPDATE,
// 并发的计费扣款不会被旧快照覆盖。
func (r *userRepository) AdjustBalance(ctx context.Context, id int64, delta float64) (service.BalanceChange, error) {
const updateSQL = `
UPDATE users
SET balance = balance + $1, updated_at = NOW()
WHERE id = $2 AND deleted_at IS NULL AND balance + $1 >= 0
RETURNING balance - $1, balance
`
change, ok, err := scanBalanceChange(ctx, clientFromContext(ctx, r.client), updateSQL, delta, id)
if err != nil {
return service.BalanceChange{}, err
}
if ok {
return change, nil
}
// 0 行既可能是用户不存在,也可能是余额不足以承受这次扣减,需要区分。
current, err := r.currentBalance(ctx, id)
if err != nil {
return service.BalanceChange{}, err
}
return service.BalanceChange{Old: current, New: current + delta}, service.ErrBalanceNegative
}
// SetBalance 原子地把余额置为 value,并返回变更前后的值。
func (r *userRepository) SetBalance(ctx context.Context, id int64, value float64) (service.BalanceChange, error) {
if value < 0 {
// 连同当前余额一起返回,便于上层给出可读的错误信息。
current, err := r.currentBalance(ctx, id)
if err != nil {
return service.BalanceChange{}, err
}
return service.BalanceChange{Old: current, New: value}, service.ErrBalanceNegative
}
const updateSQL = `
UPDATE users AS u
SET balance = $1, updated_at = NOW()
FROM (SELECT id, balance FROM users WHERE id = $2 AND deleted_at IS NULL) AS prev
WHERE u.id = prev.id AND u.deleted_at IS NULL
RETURNING prev.balance, u.balance
`
change, ok, err := scanBalanceChange(ctx, clientFromContext(ctx, r.client), updateSQL, value, id)
if err != nil {
return service.BalanceChange{}, err
}
if !ok {
return service.BalanceChange{}, service.ErrUserNotFound
}
return change, nil
}
// currentBalance 读取用户当前余额,用户不存在时返回 ErrUserNotFound。
func (r *userRepository) currentBalance(ctx context.Context, id int64) (balance float64, err error) {
rows, err := clientFromContext(ctx, r.client).QueryContext(ctx,
`SELECT balance FROM users WHERE id = $1 AND deleted_at IS NULL`, id)
if err != nil {
return 0, err
}
defer func() {
if closeErr := rows.Close(); closeErr != nil && err == nil {
err = closeErr
}
}()
if !rows.Next() {
if rowsErr := rows.Err(); rowsErr != nil {
return 0, rowsErr
}
return 0, service.ErrUserNotFound
}
if err := rows.Scan(&balance); err != nil {
return 0, err
}
return balance, rows.Err()
}
// scanBalanceChange 执行一条 RETURNING 旧余额、新余额的语句。ok 为 false 表示语句未命中任何行。
func scanBalanceChange(ctx context.Context, client *dbent.Client, query string, args ...any) (change service.BalanceChange, ok bool, err error) {
rows, err := client.QueryContext(ctx, query, args...)
if err != nil {
return service.BalanceChange{}, false, err
}
defer func() {
if closeErr := rows.Close(); closeErr != nil && err == nil {
err = closeErr
}
}()
if !rows.Next() {
if rowsErr := rows.Err(); rowsErr != nil {
return service.BalanceChange{}, false, rowsErr
}
return service.BalanceChange{}, false, nil
}
if err := rows.Scan(&change.Old, &change.New); err != nil {
return service.BalanceChange{}, false, err
}
return change, true, rows.Err()
}
func (r *userRepository) UpdateConcurrency(ctx context.Context, id int64, amount int) error {
client := clientFromContext(ctx, r.client)
n, err := client.User.Update().Where(dbuser.IDEQ(id)).AddConcurrency(amount).Save(ctx)
@@ -60,7 +60,7 @@ func (s *UserRepoSuite) TestUpdate_ReplacesEmailAuthIdentityWhenEmailChanges() {
})
user.Email = "after-update@example.com"
s.Require().NoError(s.repo.Update(s.ctx, user))
s.Require().NoError(s.repo.Update(s.ctx, user, service.UserUpdateFields{Email: true}))
newIdentity, err := s.client.AuthIdentity.Query().
Where(
@@ -118,7 +118,7 @@ func TestUserRepositoryUpdateRejectsNormalizedEmailDuplicate(t *testing.T) {
require.NoError(t, repo.Create(ctx, second))
second.Email = " existing@example.com "
err := repo.Update(ctx, second)
err := repo.Update(ctx, second, service.UserUpdateFields{Email: true})
require.ErrorIs(t, err, service.ErrEmailExists)
}
@@ -154,7 +154,7 @@ func (s *UserRepoSuite) TestUpdate() {
got, err := s.repo.GetByID(s.ctx, user.ID)
s.Require().NoError(err)
got.Username = "updated"
s.Require().NoError(s.repo.Update(s.ctx, got), "Update")
s.Require().NoError(s.repo.Update(s.ctx, got, service.UserUpdateFields{Username: true}), "Update")
updated, err := s.repo.GetByID(s.ctx, user.ID)
s.Require().NoError(err, "GetByID after update")
@@ -232,7 +232,7 @@ func (s *UserRepoSuite) TestUpdateIgnoresNoRowsFromConflictingEmailIdentityUpser
got, err := s.repo.GetByID(s.ctx, user.ID)
s.Require().NoError(err)
got.Username = "updated"
s.Require().NoError(s.repo.Update(s.ctx, got), "Update should tolerate ON CONFLICT DO NOTHING returning no rows")
s.Require().NoError(s.repo.Update(s.ctx, got, service.UserUpdateFields{Username: true}), "Update should tolerate ON CONFLICT DO NOTHING returning no rows")
updated, err := s.repo.GetByID(s.ctx, user.ID)
s.Require().NoError(err)
@@ -658,7 +658,7 @@ func (s *UserRepoSuite) TestCRUD_And_Filters_And_AtomicUpdates() {
s.Require().Equal(user2.ID, gotByEmail.ID, "GetByEmail ID mismatch")
got.Username = "Alice2"
s.Require().NoError(s.repo.Update(s.ctx, got), "Update")
s.Require().NoError(s.repo.Update(s.ctx, got, service.UserUpdateFields{Username: true}), "Update")
got2, err := s.repo.GetByID(s.ctx, user1.ID)
s.Require().NoError(err, "GetByID after update")
s.Require().Equal("Alice2", got2.Username, "Update did not persist")
@@ -0,0 +1,191 @@
//go:build integration
package repository
import (
"github.com/Wei-Shaw/sub2api/internal/service"
)
// 这一组用例覆盖用户行上的 lost update:调用方手里的快照可能早于并发发生的
// 原子写入(扣费、状态变更、限额调整、分组授予)。Update 只写显式声明的列,
// 未声明的列一律保持库中当前值,因此陈旧快照不会回滚这些并发结果。
func (s *UserRepoSuite) TestUpdate_DoesNotRevertConcurrentBalanceDeduction() {
user := s.mustCreateUser(&service.User{
Email: "lost-update-balance@example.com",
Username: "before",
Balance: 0.30,
})
// 调用方在扣费之前读到的旧快照(余额还是 0.30)。
stale, err := s.repo.GetByID(s.ctx, user.ID)
s.Require().NoError(err, "GetByID")
s.Require().InDelta(0.30, stale.Balance, 1e-9)
// 与之并发:计费按原子方式扣费。
s.Require().NoError(s.repo.DeductBalance(s.ctx, user.ID, 0.25), "DeductBalance")
// 基于旧快照的资料更新这时才落库。
stale.Username = "after"
s.Require().NoError(
s.repo.Update(s.ctx, stale, service.UserUpdateFields{Username: true}),
"Update",
)
got, err := s.repo.GetByID(s.ctx, user.ID)
s.Require().NoError(err, "GetByID after update")
s.Require().Equal("after", got.Username, "declared column must still be written")
s.Require().InDelta(0.05, got.Balance, 1e-9, "balance must not be reverted by a stale profile save")
}
// 同理,风控自动封禁把 status 置为 disabled 后,
// 基于旧快照的资料更新不得把 status 刷回 active。
func (s *UserRepoSuite) TestUpdate_DoesNotRevertConcurrentBan() {
user := s.mustCreateUser(&service.User{
Email: "lost-update-ban@example.com",
Username: "before",
Status: service.StatusActive,
})
stale, err := s.repo.GetByID(s.ctx, user.ID)
s.Require().NoError(err, "GetByID")
s.Require().Equal(service.StatusActive, stale.Status)
banned, err := s.repo.GetByID(s.ctx, user.ID)
s.Require().NoError(err, "GetByID for ban")
banned.Status = service.StatusDisabled
s.Require().NoError(
s.repo.Update(s.ctx, banned, service.UserUpdateFields{Status: true}),
"ban",
)
stale.Username = "after"
s.Require().NoError(
s.repo.Update(s.ctx, stale, service.UserUpdateFields{Username: true}),
"stale profile save",
)
got, err := s.repo.GetByID(s.ctx, user.ID)
s.Require().NoError(err, "GetByID after update")
s.Require().Equal("after", got.Username)
s.Require().Equal(service.StatusDisabled, got.Status, "ban must survive a stale profile save")
}
// 未声明的列不写,也意味着并发的限额调整不会被资料保存回滚。
func (s *UserRepoSuite) TestUpdate_DoesNotRevertConcurrentLimitChanges() {
user := s.mustCreateUser(&service.User{
Email: "lost-update-limits@example.com",
Username: "before",
Concurrency: 3,
RPMLimit: 30,
})
stale, err := s.repo.GetByID(s.ctx, user.ID)
s.Require().NoError(err, "GetByID")
concurrency, rpmLimit := 9, 90
affected, err := s.repo.BatchUpdateLimits(s.ctx, []int64{user.ID}, &concurrency, &rpmLimit)
s.Require().NoError(err, "BatchUpdateLimits")
s.Require().Equal(1, affected)
stale.Username = "after"
s.Require().NoError(
s.repo.Update(s.ctx, stale, service.UserUpdateFields{Username: true}),
"stale profile save",
)
got, err := s.repo.GetByID(s.ctx, user.ID)
s.Require().NoError(err, "GetByID after update")
s.Require().Equal(9, got.Concurrency, "concurrency must not be reverted")
s.Require().Equal(90, got.RPMLimit, "rpm limit must not be reverted")
}
// AllowedGroups 只在显式声明时才同步,否则并发授予的分组权限会被旧快照删掉。
func (s *UserRepoSuite) TestUpdate_DoesNotRevertConcurrentAllowedGroupGrant() {
group := s.mustCreateGroup("lost-update-group")
user := s.mustCreateUser(&service.User{
Email: "lost-update-groups@example.com",
Username: "before",
})
stale, err := s.repo.GetByID(s.ctx, user.ID)
s.Require().NoError(err, "GetByID")
s.Require().Empty(stale.AllowedGroups)
s.Require().NoError(s.repo.AddGroupToAllowedGroups(s.ctx, user.ID, group.ID), "AddGroupToAllowedGroups")
stale.Username = "after"
s.Require().NoError(
s.repo.Update(s.ctx, stale, service.UserUpdateFields{Username: true}),
"stale profile save",
)
got, err := s.repo.GetByID(s.ctx, user.ID)
s.Require().NoError(err, "GetByID after update")
s.Require().Equal([]int64{group.ID}, got.AllowedGroups, "granted group must not be reverted")
}
func (s *UserRepoSuite) TestAdjustBalance_AppliesDeltaAndReportsChange() {
user := s.mustCreateUser(&service.User{Email: "adjust-balance@example.com", Balance: 10})
change, err := s.repo.AdjustBalance(s.ctx, user.ID, 5)
s.Require().NoError(err, "AdjustBalance add")
s.Require().InDelta(10, change.Old, 1e-9)
s.Require().InDelta(15, change.New, 1e-9)
change, err = s.repo.AdjustBalance(s.ctx, user.ID, -5)
s.Require().NoError(err, "AdjustBalance subtract")
s.Require().InDelta(15, change.Old, 1e-9)
s.Require().InDelta(10, change.New, 1e-9)
got, err := s.repo.GetByID(s.ctx, user.ID)
s.Require().NoError(err, "GetByID")
s.Require().InDelta(10, got.Balance, 1e-9)
}
func (s *UserRepoSuite) TestAdjustBalance_RefusesNegativeResult() {
user := s.mustCreateUser(&service.User{Email: "adjust-balance-negative@example.com", Balance: 3})
change, err := s.repo.AdjustBalance(s.ctx, user.ID, -4)
s.Require().ErrorIs(err, service.ErrBalanceNegative)
s.Require().InDelta(3, change.Old, 1e-9, "error must report the real current balance")
s.Require().InDelta(-1, change.New, 1e-9)
got, err := s.repo.GetByID(s.ctx, user.ID)
s.Require().NoError(err, "GetByID")
s.Require().InDelta(3, got.Balance, 1e-9, "refused adjustment must not write")
}
func (s *UserRepoSuite) TestAdjustBalance_UserNotFound() {
_, err := s.repo.AdjustBalance(s.ctx, 99999999, 1)
s.Require().ErrorIs(err, service.ErrUserNotFound)
}
func (s *UserRepoSuite) TestSetBalance_ReplacesValueAndReportsPrevious() {
user := s.mustCreateUser(&service.User{Email: "set-balance@example.com", Balance: 7})
change, err := s.repo.SetBalance(s.ctx, user.ID, 2)
s.Require().NoError(err, "SetBalance")
s.Require().InDelta(7, change.Old, 1e-9)
s.Require().InDelta(2, change.New, 1e-9)
got, err := s.repo.GetByID(s.ctx, user.ID)
s.Require().NoError(err, "GetByID")
s.Require().InDelta(2, got.Balance, 1e-9)
}
func (s *UserRepoSuite) TestSetBalance_RejectsNegativeValue() {
user := s.mustCreateUser(&service.User{Email: "set-balance-negative@example.com", Balance: 7})
_, err := s.repo.SetBalance(s.ctx, user.ID, -1)
s.Require().ErrorIs(err, service.ErrBalanceNegative)
got, err := s.repo.GetByID(s.ctx, user.ID)
s.Require().NoError(err, "GetByID")
s.Require().InDelta(7, got.Balance, 1e-9)
}
func (s *UserRepoSuite) TestSetBalance_UserNotFound() {
_, err := s.repo.SetBalance(s.ctx, 99999999, 1)
s.Require().ErrorIs(err, service.ErrUserNotFound)
}
@@ -84,7 +84,7 @@ func (s *UserRepoSuite) TestUpdate_PersistsSignupSourceAndActivityTimestamps() {
created.LastLoginAt = &lastLoginAt
created.LastActiveAt = &lastActiveAt
s.Require().NoError(s.repo.Update(s.ctx, created))
s.Require().NoError(s.repo.Update(s.ctx, created, service.UserUpdateFields{SignupSource: true, LastLoginAt: true, LastActiveAt: true}))
got, err := s.repo.GetByID(s.ctx, created.ID)
s.Require().NoError(err)
+10 -2
View File
@@ -1536,7 +1536,7 @@ func (r *stubUserRepo) GetFirstAdmin(ctx context.Context) (*service.User, error)
return nil, service.ErrUserNotFound
}
func (r *stubUserRepo) Update(ctx context.Context, user *service.User) error {
func (r *stubUserRepo) Update(ctx context.Context, user *service.User, fields service.UserUpdateFields) error {
return errors.New("not implemented")
}
@@ -1572,6 +1572,14 @@ func (r *stubUserRepo) DeductBalance(ctx context.Context, id int64, amount float
return errors.New("not implemented")
}
func (r *stubUserRepo) AdjustBalance(ctx context.Context, id int64, delta float64) (service.BalanceChange, error) {
return service.BalanceChange{}, errors.New("not implemented")
}
func (r *stubUserRepo) SetBalance(ctx context.Context, id int64, value float64) (service.BalanceChange, error) {
return service.BalanceChange{}, errors.New("not implemented")
}
func (r *stubUserRepo) UpdateConcurrency(ctx context.Context, id int64, amount int) error {
return errors.New("not implemented")
}
@@ -2282,7 +2290,7 @@ func (r *stubApiKeyRepo) GetByKeyForAuth(ctx context.Context, key string) (*serv
return r.GetByKey(ctx, key)
}
func (r *stubApiKeyRepo) Update(ctx context.Context, key *service.APIKey) error {
func (r *stubApiKeyRepo) Update(ctx context.Context, key *service.APIKey, _ service.APIKeyUpdateFields) error {
if key == nil {
return errors.New("nil key")
}
@@ -150,7 +150,7 @@ func (s *stubUserRepo) GetFirstAdmin(ctx context.Context) (*service.User, error)
panic("unexpected GetFirstAdmin call")
}
func (s *stubUserRepo) Update(ctx context.Context, user *service.User) error {
func (s *stubUserRepo) Update(ctx context.Context, user *service.User, fields service.UserUpdateFields) error {
panic("unexpected Update call")
}
@@ -198,6 +198,14 @@ func (s *stubUserRepo) DeductBalance(ctx context.Context, id int64, amount float
panic("unexpected DeductBalance call")
}
func (s *stubUserRepo) AdjustBalance(ctx context.Context, id int64, delta float64) (service.BalanceChange, error) {
panic("unexpected AdjustBalance call")
}
func (s *stubUserRepo) SetBalance(ctx context.Context, id int64, value float64) (service.BalanceChange, error) {
panic("unexpected SetBalance call")
}
func (s *stubUserRepo) UpdateConcurrency(ctx context.Context, id int64, amount int) error {
panic("unexpected UpdateConcurrency call")
}
@@ -106,7 +106,7 @@ func (f fakeAPIKeyRepo) GetByKey(ctx context.Context, key string) (*service.APIK
func (f fakeAPIKeyRepo) GetByKeyForAuth(ctx context.Context, key string) (*service.APIKey, error) {
return f.GetByKey(ctx, key)
}
func (f fakeAPIKeyRepo) Update(ctx context.Context, key *service.APIKey) error {
func (f fakeAPIKeyRepo) Update(ctx context.Context, key *service.APIKey, _ service.APIKeyUpdateFields) error {
return errors.New("not implemented")
}
func (f fakeAPIKeyRepo) Delete(ctx context.Context, id int64) error {
@@ -1550,7 +1550,7 @@ func (r *stubApiKeyRepo) GetByKeyForAuth(ctx context.Context, key string) (*serv
return r.GetByKey(ctx, key)
}
func (r *stubApiKeyRepo) Update(ctx context.Context, key *service.APIKey) error {
func (r *stubApiKeyRepo) Update(ctx context.Context, key *service.APIKey, _ service.APIKeyUpdateFields) error {
return errors.New("not implemented")
}
+3 -3
View File
@@ -1078,7 +1078,7 @@ func (s *adminServiceImpl) AdminUpdateAPIKeyGroupID(ctx context.Context, keyID i
if addErr := s.userRepo.AddGroupToAllowedGroups(opCtx, apiKey.UserID, gid); addErr != nil {
return nil, fmt.Errorf("add group to user allowed groups: %w", addErr)
}
if err := s.apiKeyRepo.Update(opCtx, apiKey); err != nil {
if err := s.apiKeyRepo.Update(opCtx, apiKey, APIKeyUpdateFields{GroupID: true}); err != nil {
return nil, fmt.Errorf("update api key: %w", err)
}
if tx != nil {
@@ -1102,7 +1102,7 @@ func (s *adminServiceImpl) AdminUpdateAPIKeyGroupID(ctx context.Context, keyID i
}
// 非专属分组 / 解绑:无需事务,单步更新即可
if err := s.apiKeyRepo.Update(ctx, apiKey); err != nil {
if err := s.apiKeyRepo.Update(ctx, apiKey, APIKeyUpdateFields{GroupID: true}); err != nil {
return nil, fmt.Errorf("update api key: %w", err)
}
@@ -1127,7 +1127,7 @@ func (s *adminServiceImpl) AdminResetAPIKeyRateLimitUsage(ctx context.Context, k
apiKey.Window5hStart = nil
apiKey.Window1dStart = nil
apiKey.Window7dStart = nil
if err := s.apiKeyRepo.Update(ctx, apiKey); err != nil {
if err := s.apiKeyRepo.Update(ctx, apiKey, APIKeyUpdateFields{RateLimitUsage: true}); err != nil {
return nil, fmt.Errorf("reset api key rate limit usage: %w", err)
}
if s.authCacheInvalidator != nil {
@@ -45,7 +45,9 @@ func (s *userRepoStubForGroupUpdate) GetByEmail(context.Context, string) (*User,
func (s *userRepoStubForGroupUpdate) GetFirstAdmin(context.Context) (*User, error) {
panic("unexpected")
}
func (s *userRepoStubForGroupUpdate) Update(context.Context, *User) error { panic("unexpected") }
func (s *userRepoStubForGroupUpdate) Update(context.Context, *User, UserUpdateFields) error {
panic("unexpected")
}
func (s *userRepoStubForGroupUpdate) Delete(context.Context, int64) error { panic("unexpected") }
func (s *userRepoStubForGroupUpdate) GetUserAvatar(context.Context, int64) (*UserAvatar, error) {
panic("unexpected")
@@ -68,6 +70,14 @@ func (s *userRepoStubForGroupUpdate) UpdateBalance(context.Context, int64, float
func (s *userRepoStubForGroupUpdate) DeductBalance(context.Context, int64, float64) error {
panic("unexpected")
}
func (s *userRepoStubForGroupUpdate) AdjustBalance(ctx context.Context, id int64, delta float64) (BalanceChange, error) {
panic("unexpected AdjustBalance call")
}
func (s *userRepoStubForGroupUpdate) SetBalance(ctx context.Context, id int64, value float64) (BalanceChange, error) {
panic("unexpected SetBalance call")
}
func (s *userRepoStubForGroupUpdate) UpdateConcurrency(context.Context, int64, int) error {
panic("unexpected")
}
@@ -134,7 +144,7 @@ func (s *apiKeyRepoStubForGroupUpdate) GetByID(_ context.Context, _ int64) (*API
clone := *s.key
return &clone, nil
}
func (s *apiKeyRepoStubForGroupUpdate) Update(_ context.Context, key *APIKey) error {
func (s *apiKeyRepoStubForGroupUpdate) Update(_ context.Context, key *APIKey, _ APIKeyUpdateFields) error {
if s.updateErr != nil {
return s.updateErr
}
@@ -86,7 +86,7 @@ func (s *userRepoStub) GetFirstAdmin(ctx context.Context) (*User, error) {
panic("unexpected GetFirstAdmin call")
}
func (s *userRepoStub) Update(ctx context.Context, user *User) error {
func (s *userRepoStub) Update(ctx context.Context, user *User, fields UserUpdateFields) error {
s.updated = append(s.updated, user)
if s.usersByEmail == nil {
s.usersByEmail = make(map[string]*User)
@@ -141,6 +141,14 @@ func (s *userRepoStub) DeductBalance(ctx context.Context, id int64, amount float
panic("unexpected DeductBalance call")
}
func (s *userRepoStub) AdjustBalance(ctx context.Context, id int64, delta float64) (BalanceChange, error) {
panic("unexpected AdjustBalance call")
}
func (s *userRepoStub) SetBalance(ctx context.Context, id int64, value float64) (BalanceChange, error) {
panic("unexpected SetBalance call")
}
func (s *userRepoStub) UpdateConcurrency(ctx context.Context, id int64, amount int) error {
panic("unexpected UpdateConcurrency call")
}
@@ -64,7 +64,7 @@ func (s *emailSyncRepoStub) GetFirstAdmin(context.Context) (*User, error) {
return nil, fmt.Errorf("unexpected GetFirstAdmin call")
}
func (s *emailSyncRepoStub) Update(_ context.Context, user *User) error {
func (s *emailSyncRepoStub) Update(_ context.Context, user *User, _ UserUpdateFields) error {
s.updateCalls++
s.updated = append(s.updated, user)
s.user = user
@@ -117,6 +117,14 @@ func (s *emailSyncRepoStub) ExistsByEmailAlias(context.Context, string) (bool, e
return false, nil
}
func (s *emailSyncRepoStub) AdjustBalance(ctx context.Context, id int64, delta float64) (BalanceChange, error) {
panic("unexpected AdjustBalance call")
}
func (s *emailSyncRepoStub) SetBalance(ctx context.Context, id int64, value float64) (BalanceChange, error) {
panic("unexpected SetBalance call")
}
func (s *emailSyncRepoStub) RemoveGroupFromAllowedGroups(context.Context, int64) (int64, error) {
return 0, nil
}
@@ -12,23 +12,34 @@ import (
type balanceUserRepoStub struct {
*userRepoStub
updateErr error
updated []*User
adjustErr error
// changes 记录每次原子余额变更,顺序与调用顺序一致。
changes []BalanceChange
}
func (s *balanceUserRepoStub) Update(ctx context.Context, user *User) error {
if s.updateErr != nil {
return s.updateErr
func (s *balanceUserRepoStub) AdjustBalance(ctx context.Context, id int64, delta float64) (BalanceChange, error) {
return s.apply(func(current float64) float64 { return current + delta })
}
func (s *balanceUserRepoStub) SetBalance(ctx context.Context, id int64, value float64) (BalanceChange, error) {
return s.apply(func(float64) float64 { return value })
}
func (s *balanceUserRepoStub) apply(next func(current float64) float64) (BalanceChange, error) {
if s.adjustErr != nil {
return BalanceChange{}, s.adjustErr
}
if user == nil {
return nil
if s.userRepoStub == nil || s.userRepoStub.user == nil {
return BalanceChange{}, ErrUserNotFound
}
clone := *user
s.updated = append(s.updated, &clone)
if s.userRepoStub != nil {
s.userRepoStub.user = &clone
change := BalanceChange{Old: s.userRepoStub.user.Balance}
change.New = next(change.Old)
if change.New < 0 {
return change, ErrBalanceNegative
}
return nil
s.userRepoStub.user.Balance = change.New
s.changes = append(s.changes, change)
return change, nil
}
type balanceRedeemRepoStub struct {
@@ -87,6 +98,63 @@ func (s *authCacheInvalidatorStub) InvalidateAuthCacheByGroupID(ctx context.Cont
s.groupIDs = append(s.groupIDs, groupID)
}
// 管理员调账必须走原子的 AdjustBalance/SetBalance,而不是"读余额→算新值→整行写回",
// 后者会把并发的计费扣款覆盖掉。userRepoStub.Update 对未预期的调用会 panic,
// 因此这里同时证明它没被走到。
func TestAdminService_UpdateUserBalance_UsesAtomicPrimitives(t *testing.T) {
tests := []struct {
name string
operation string
amount float64
want BalanceChange
}{
{name: "add", operation: "add", amount: 5, want: BalanceChange{Old: 10, New: 15}},
{name: "subtract", operation: "subtract", amount: 4, want: BalanceChange{Old: 10, New: 6}},
{name: "set", operation: "set", amount: 2, want: BalanceChange{Old: 10, New: 2}},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
repo := &balanceUserRepoStub{userRepoStub: &userRepoStub{user: &User{ID: 7, Balance: 10}}}
svc := &adminServiceImpl{
userRepo: repo,
redeemCodeRepo: &balanceRedeemRepoStub{redeemRepoStub: &redeemRepoStub{}},
}
user, err := svc.UpdateUserBalance(context.Background(), 7, tt.amount, tt.operation, "")
require.NoError(t, err)
require.Equal(t, []BalanceChange{tt.want}, repo.changes)
require.Equal(t, tt.want.New, user.Balance)
})
}
}
func TestAdminService_UpdateUserBalance_RejectsNegativeResult(t *testing.T) {
repo := &balanceUserRepoStub{userRepoStub: &userRepoStub{user: &User{ID: 7, Balance: 3}}}
svc := &adminServiceImpl{
userRepo: repo,
redeemCodeRepo: &balanceRedeemRepoStub{redeemRepoStub: &redeemRepoStub{}},
}
_, err := svc.UpdateUserBalance(context.Background(), 7, 4, "subtract", "")
require.Error(t, err)
require.Contains(t, err.Error(), "balance cannot be negative")
require.Empty(t, repo.changes, "refused adjustment must not be applied")
require.Equal(t, 3.0, repo.userRepoStub.user.Balance)
}
func TestAdminService_UpdateUserBalance_RejectsUnknownOperation(t *testing.T) {
repo := &balanceUserRepoStub{userRepoStub: &userRepoStub{user: &User{ID: 7, Balance: 10}}}
svc := &adminServiceImpl{
userRepo: repo,
redeemCodeRepo: &balanceRedeemRepoStub{redeemRepoStub: &redeemRepoStub{}},
}
_, err := svc.UpdateUserBalance(context.Background(), 7, 1, "multiply", "")
require.Error(t, err)
require.Empty(t, repo.changes)
}
func TestAdminService_UpdateUserBalance_InvalidatesAuthCache(t *testing.T) {
baseRepo := &userRepoStub{user: &User{ID: 7, Balance: 10}}
repo := &balanceUserRepoStub{userRepoStub: baseRepo}
@@ -16,7 +16,7 @@ type rpmUserRepoStub struct {
lastUpdated *User
}
func (s *rpmUserRepoStub) Update(_ context.Context, user *User) error {
func (s *rpmUserRepoStub) Update(_ context.Context, user *User, _ UserUpdateFields) error {
if user == nil {
return nil
}
+37 -20
View File
@@ -218,24 +218,33 @@ func (s *adminServiceImpl) UpdateUser(ctx context.Context, id int64, input *Upda
oldRPMLimit := user.RPMLimit
oldAllowedGroups := append([]int64(nil), user.AllowedGroups...)
// fields 与下面的 input.X 判空条件一一对应:管理员没提交的列不写回,
// 避免这份快照回滚并发的扣费、状态变更或批量限额调整。
var fields UserUpdateFields
if input.Email != "" {
user.Email = input.Email
fields.Email = true
}
if input.Password != "" {
if err := user.SetPassword(input.Password); err != nil {
return nil, err
}
fields.PasswordHash = true
}
if input.Username != nil {
user.Username = *input.Username
fields.Username = true
}
if input.Notes != nil {
user.Notes = *input.Notes
fields.Notes = true
}
if input.Status != "" {
user.Status = input.Status
fields.Status = true
}
// 角色变更(admin/user);空字符串表示不修改。
@@ -252,21 +261,25 @@ func (s *adminServiceImpl) UpdateUser(ctx context.Context, id int64, input *Upda
}
}
user.Role = role
fields.Role = true
}
if input.Concurrency != nil {
user.Concurrency = *input.Concurrency
fields.Concurrency = true
}
if input.RPMLimit != nil {
user.RPMLimit = *input.RPMLimit
fields.RPMLimit = true
}
if input.AllowedGroups != nil {
user.AllowedGroups = *input.AllowedGroups
fields.AllowedGroups = true
}
if err := s.userRepo.Update(ctx, user); err != nil {
if err := s.userRepo.Update(ctx, user, fields); err != nil {
return nil, err
}
@@ -493,30 +506,34 @@ func (s *adminServiceImpl) BatchUpdateLimits(ctx context.Context, userIDs []int6
}
func (s *adminServiceImpl) UpdateUserBalance(ctx context.Context, userID int64, balance float64, operation string, notes string) (*User, error) {
// 余额调整必须走原子接口:先读后整行写回会把并发的计费扣款覆盖掉。
var (
change BalanceChange
err error
)
switch operation {
case "set":
change, err = s.userRepo.SetBalance(ctx, userID, balance)
case "add":
change, err = s.userRepo.AdjustBalance(ctx, userID, balance)
case "subtract":
change, err = s.userRepo.AdjustBalance(ctx, userID, -balance)
default:
return nil, fmt.Errorf("unsupported balance operation: %q", operation)
}
if errors.Is(err, ErrBalanceNegative) {
return nil, fmt.Errorf("balance cannot be negative, current balance: %.2f, requested operation would result in: %.2f", change.Old, change.New)
}
if err != nil {
return nil, err
}
user, err := s.userRepo.GetByID(ctx, userID)
if err != nil {
return nil, err
}
oldBalance := user.Balance
switch operation {
case "set":
user.Balance = balance
case "add":
user.Balance += balance
case "subtract":
user.Balance -= balance
}
if user.Balance < 0 {
return nil, fmt.Errorf("balance cannot be negative, current balance: %.2f, requested operation would result in: %.2f", oldBalance, user.Balance)
}
if err := s.userRepo.Update(ctx, user); err != nil {
return nil, err
}
balanceDiff := user.Balance - oldBalance
balanceDiff := change.New - change.Old
if s.authCacheInvalidator != nil && balanceDiff != 0 {
s.authCacheInvalidator.InvalidateAuthCacheByUserID(ctx, userID)
}
+59 -3
View File
@@ -53,6 +53,34 @@ const (
apiKeyLastUsedFailBackoff = 5 * time.Second
)
// APIKeyUpdateFields 声明 APIKeyRepository.Update 允许写回的列。
//
// 与 UserUpdateFields 同理:api_keys 的用量列由计费热路径原子递增
// (IncrementQuotaUsed / IncrementRateLimitUsage 的 quota_used、usage_5h/1d/7d),
// 若编辑 Key 时无条件整行回写,并发累计的配额与限流计数就会被旧快照覆盖。
// 因此调用方必须显式声明要改的列。
type APIKeyUpdateFields struct {
Name bool
Status bool
Quota bool
GroupID bool
ExpiresAt bool
// QuotaUsed 仅供"重置配额用量"路径声明;常规计费走 IncrementQuotaUsed。
QuotaUsed bool
// RateLimits 覆盖 rate_limit_5h / _1d / _7d 三个阈值。
RateLimits bool
// RateLimitUsage 覆盖 usage_5h/_1d/_7d 与三个窗口起点,
// 仅供"重置限流用量"路径声明;常规计费走 IncrementRateLimitUsage。
RateLimitUsage bool
// IPRules 覆盖 ip_whitelist 与 ip_blacklist。
IPRules bool
}
// IsEmpty 报告该次 Update 是否不写任何列。
func (f APIKeyUpdateFields) IsEmpty() bool {
return f == APIKeyUpdateFields{}
}
type APIKeyRepository interface {
Create(ctx context.Context, key *APIKey) error
GetByID(ctx context.Context, id int64) (*APIKey, error)
@@ -61,7 +89,8 @@ type APIKeyRepository interface {
GetByKey(ctx context.Context, key string) (*APIKey, error)
// GetByKeyForAuth 认证专用查询,返回最小字段集
GetByKeyForAuth(ctx context.Context, key string) (*APIKey, error)
Update(ctx context.Context, key *APIKey) error
// Update 只写 fields 中显式声明的列,其余列保持库中当前值。
Update(ctx context.Context, key *APIKey, fields APIKeyUpdateFields) error
Delete(ctx context.Context, id int64) error
// DeleteWithAudit keeps the legacy interface name for rolling-upgrade compatibility.
// Implementations must tombstone the key and soft-delete it atomically without
@@ -718,9 +747,17 @@ func (s *APIKeyService) Update(ctx context.Context, id int64, userID int64, req
}
}
// fields 只登记本次请求真正要改的列。quota_used 与 usage_5h/1d/7d 由计费热路径
// 原子递增,除非用户显式点了"重置",否则这里不用快照把它们写回去。
var fields APIKeyUpdateFields
// 下面若干分支会顺带把 Status 改回 active(配额扩容、清除过期等),
// 所以用原始值比对来决定是否写 status,而不是只看 req.Status。
originalStatus := apiKey.Status
// 更新字段
if req.Name != nil {
apiKey.Name = html.EscapeString(*req.Name)
fields.Name = true
}
if req.GroupID != nil {
@@ -740,10 +777,12 @@ func (s *APIKeyService) Update(ctx context.Context, id int64, userID int64, req
}
apiKey.GroupID = req.GroupID
fields.GroupID = true
}
if req.Status != nil {
apiKey.Status = *req.Status
fields.Status = true
// 如果状态改变,清除Redis缓存
if s.cache != nil {
_ = s.cache.DeleteCreateAttemptCount(ctx, apiKey.UserID)
@@ -753,6 +792,7 @@ func (s *APIKeyService) Update(ctx context.Context, id int64, userID int64, req
// Update quota fields
if req.Quota != nil {
apiKey.Quota = *req.Quota
fields.Quota = true
// If quota now has room, or is changed to unlimited, reactivate exhausted keys.
if apiKey.Status == StatusAPIKeyQuotaExhausted && (*req.Quota <= 0 || *req.Quota > apiKey.QuotaUsed) {
apiKey.Status = StatusActive
@@ -760,6 +800,7 @@ func (s *APIKeyService) Update(ctx context.Context, id int64, userID int64, req
}
if req.ResetQuota != nil && *req.ResetQuota {
apiKey.QuotaUsed = 0
fields.QuotaUsed = true
// If resetting quota and status was quota_exhausted, reactivate
if apiKey.Status == StatusAPIKeyQuotaExhausted {
apiKey.Status = StatusActive
@@ -767,12 +808,14 @@ func (s *APIKeyService) Update(ctx context.Context, id int64, userID int64, req
}
if req.ClearExpiration {
apiKey.ExpiresAt = nil
fields.ExpiresAt = true
// If clearing expiry and status was expired, reactivate
if apiKey.Status == StatusAPIKeyExpired {
apiKey.Status = StatusActive
}
} else if req.ExpiresAt != nil {
apiKey.ExpiresAt = req.ExpiresAt
fields.ExpiresAt = true
// If extending expiry and status was expired, reactivate
if apiKey.Status == StatusAPIKeyExpired && time.Now().Before(*req.ExpiresAt) {
apiKey.Status = StatusActive
@@ -782,20 +825,25 @@ func (s *APIKeyService) Update(ctx context.Context, id int64, userID int64, req
// 更新 IP 限制(nil 不修改,空数组清空设置)
if req.IPWhitelist != nil {
apiKey.IPWhitelist = *req.IPWhitelist
fields.IPRules = true
}
if req.IPBlacklist != nil {
apiKey.IPBlacklist = *req.IPBlacklist
fields.IPRules = true
}
// Update rate limit configuration
if req.RateLimit5h != nil {
apiKey.RateLimit5h = *req.RateLimit5h
fields.RateLimits = true
}
if req.RateLimit1d != nil {
apiKey.RateLimit1d = *req.RateLimit1d
fields.RateLimits = true
}
if req.RateLimit7d != nil {
apiKey.RateLimit7d = *req.RateLimit7d
fields.RateLimits = true
}
resetRateLimit := req.ResetRateLimitUsage != nil && *req.ResetRateLimitUsage
if resetRateLimit {
@@ -805,9 +853,15 @@ func (s *APIKeyService) Update(ctx context.Context, id int64, userID int64, req
apiKey.Window5hStart = nil
apiKey.Window1dStart = nil
apiKey.Window7dStart = nil
fields.RateLimitUsage = true
}
if err := s.apiKeyRepo.Update(ctx, apiKey); err != nil {
// 上面的自动复活分支可能改了 status,这里统一登记。
if apiKey.Status != originalStatus {
fields.Status = true
}
if err := s.apiKeyRepo.Update(ctx, apiKey, fields); err != nil {
return nil, fmt.Errorf("update api key: %w", err)
}
@@ -1046,7 +1100,9 @@ func (s *APIKeyService) UpdateQuotaUsed(ctx context.Context, apiKeyID int64, cos
// If quota is set and now exhausted, update status
if apiKey.Quota > 0 && newQuotaUsed >= apiKey.Quota {
apiKey.Status = StatusAPIKeyQuotaExhausted
if err := s.apiKeyRepo.Update(ctx, apiKey); err != nil {
// 只写 status:这条位于计费热路径,若整行回写会把刚刚原子递增的
// quota_used 与限流用量按快照覆盖掉。
if err := s.apiKeyRepo.Update(ctx, apiKey, APIKeyUpdateFields{Status: true}); err != nil {
return nil // Don't fail the request
}
// Invalidate cache so next request sees the new status
@@ -46,7 +46,7 @@ func (s *authRepoStub) GetByKeyForAuth(ctx context.Context, key string) (*APIKey
return s.getByKeyForAuth(ctx, key)
}
func (s *authRepoStub) Update(ctx context.Context, key *APIKey) error {
func (s *authRepoStub) Update(ctx context.Context, key *APIKey, _ APIKeyUpdateFields) error {
panic("unexpected Update call")
}
@@ -82,7 +82,7 @@ func (s *apiKeyRepoStub) GetByKeyForAuth(ctx context.Context, key string) (*APIK
panic("unexpected GetByKeyForAuth call")
}
func (s *apiKeyRepoStub) Update(ctx context.Context, key *APIKey) error {
func (s *apiKeyRepoStub) Update(ctx context.Context, key *APIKey, _ APIKeyUpdateFields) error {
if key != nil {
s.updatedKeys = append(s.updatedKeys, *key)
}
@@ -95,7 +95,7 @@ func (s *quotaBaseAPIKeyRepoStub) GetByKey(context.Context, string) (*APIKey, er
func (s *quotaBaseAPIKeyRepoStub) GetByKeyForAuth(context.Context, string) (*APIKey, error) {
panic("unexpected GetByKeyForAuth call")
}
func (s *quotaBaseAPIKeyRepoStub) Update(context.Context, *APIKey) error {
func (s *quotaBaseAPIKeyRepoStub) Update(context.Context, *APIKey, APIKeyUpdateFields) error {
panic("unexpected Update call")
}
func (s *quotaBaseAPIKeyRepoStub) Delete(context.Context, int64) error {
@@ -0,0 +1,133 @@
//go:build unit
package service
import (
"context"
"testing"
"github.com/stretchr/testify/require"
)
// api_keys 的 quota_used / usage_5h|1d|7d 由计费热路径原子递增。
// 编辑 Key(改名、换分组……)若整行回写,并发累计的用量就会被旧快照覆盖。
// 这些用例锁死"只声明请求真正要改的列"。
type updateFieldsAPIKeyRepoStub struct {
quotaBaseAPIKeyRepoStub
key *APIKey
updateFields []APIKeyUpdateFields
}
// IncrementQuotaUsed 模拟计费热路径上的原子递增:只动 quota_used。
func (s *updateFieldsAPIKeyRepoStub) IncrementQuotaUsed(_ context.Context, _ int64, amount float64) (float64, error) {
s.key.QuotaUsed += amount
return s.key.QuotaUsed, nil
}
func (s *updateFieldsAPIKeyRepoStub) GetByID(context.Context, int64) (*APIKey, error) {
clone := *s.key
return &clone, nil
}
func (s *updateFieldsAPIKeyRepoStub) Update(_ context.Context, _ *APIKey, fields APIKeyUpdateFields) error {
s.updateFields = append(s.updateFields, fields)
return nil
}
func newUpdateFieldsAPIKeyService(key *APIKey) (*APIKeyService, *updateFieldsAPIKeyRepoStub) {
repo := &updateFieldsAPIKeyRepoStub{key: key}
return &APIKeyService{apiKeyRepo: repo}, repo
}
func TestAPIKeyUpdate_OnlyDeclaresRequestedColumns(t *testing.T) {
name := "renamed"
quota := 500.0
rateLimit := 42.0
whitelist := []string{"10.0.0.1"}
tests := []struct {
name string
req UpdateAPIKeyRequest
want APIKeyUpdateFields
}{
{
name: "name only",
req: UpdateAPIKeyRequest{Name: &name},
want: APIKeyUpdateFields{Name: true},
},
{
name: "quota only",
req: UpdateAPIKeyRequest{Quota: &quota},
want: APIKeyUpdateFields{Quota: true},
},
{
name: "rate limit threshold only",
req: UpdateAPIKeyRequest{RateLimit5h: &rateLimit},
want: APIKeyUpdateFields{RateLimits: true},
},
{
name: "ip whitelist only",
req: UpdateAPIKeyRequest{IPWhitelist: &whitelist},
want: APIKeyUpdateFields{IPRules: true},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
svc, repo := newUpdateFieldsAPIKeyService(&APIKey{
ID: 1,
UserID: 7,
Key: "sk-test",
Name: "before",
Status: StatusActive,
Quota: 100,
QuotaUsed: 30,
Usage5h: 12,
})
_, err := svc.Update(context.Background(), 1, 7, tt.req)
require.NoError(t, err)
require.Equal(t, []APIKeyUpdateFields{tt.want}, repo.updateFields)
})
}
}
// 显式重置仍需声明对应的列,避免收窄写入列时把功能改坏。
func TestAPIKeyUpdate_DeclaresUsageColumnsOnExplicitReset(t *testing.T) {
reset := true
svc, repo := newUpdateFieldsAPIKeyService(&APIKey{
ID: 1, UserID: 7, Key: "sk-test", Status: StatusActive, Quota: 100, QuotaUsed: 30, Usage5h: 12,
})
_, err := svc.Update(context.Background(), 1, 7, UpdateAPIKeyRequest{
ResetQuota: &reset,
ResetRateLimitUsage: &reset,
})
require.NoError(t, err)
require.Equal(t, []APIKeyUpdateFields{{QuotaUsed: true, RateLimitUsage: true}}, repo.updateFields)
}
// 配额扩容会顺带把 quota_exhausted 复活为 active,此时必须声明 status。
func TestAPIKeyUpdate_DeclaresStatusWhenReactivated(t *testing.T) {
quota := 500.0
svc, repo := newUpdateFieldsAPIKeyService(&APIKey{
ID: 1, UserID: 7, Key: "sk-test", Status: StatusAPIKeyQuotaExhausted, Quota: 100, QuotaUsed: 100,
})
_, err := svc.Update(context.Background(), 1, 7, UpdateAPIKeyRequest{Quota: &quota})
require.NoError(t, err)
require.Equal(t, []APIKeyUpdateFields{{Quota: true, Status: true}}, repo.updateFields)
}
// 计费热路径把 Key 标记为配额耗尽时只写 status,
// 否则会把刚原子递增的 quota_used 按快照覆盖掉。
func TestUpdateQuotaUsed_ExhaustedMarkOnlyDeclaresStatus(t *testing.T) {
repo := &updateFieldsAPIKeyRepoStub{key: &APIKey{
ID: 1, UserID: 7, Key: "sk-test", Status: StatusActive, Quota: 10, QuotaUsed: 10,
}}
svc := &APIKeyService{apiKeyRepo: repo}
require.NoError(t, svc.UpdateQuotaUsed(context.Background(), 1, 5))
require.Equal(t, []APIKeyUpdateFields{{Status: true}}, repo.updateFields)
}
@@ -79,7 +79,7 @@ func (s *AuthService) BindEmailIdentity(
currentUser.Email = normalizedEmail
currentUser.PasswordHash = hashedPassword
if err := s.userRepo.Update(ctx, currentUser); err != nil {
if err := s.userRepo.Update(ctx, currentUser, UserUpdateFields{Email: true, PasswordHash: true}); err != nil {
if errors.Is(err, ErrEmailExists) {
return nil, ErrEmailExists
}
@@ -134,7 +134,7 @@ func (s *AuthService) loginOrRegisterVerifiedEmailOAuth(
if user.Username == "" && strings.TrimSpace(input.Username) != "" {
user.Username = strings.TrimSpace(input.Username)
if err := s.userRepo.Update(ctx, user); err != nil {
if err := s.userRepo.Update(ctx, user, UserUpdateFields{Username: true}); err != nil {
logger.LegacyPrintf("service.auth", "[Auth] Failed to update username after %s oauth login: %v", providerType, err)
}
}
+12 -12
View File
@@ -561,7 +561,7 @@ func (s *AuthService) LoginOrRegisterOAuth(ctx context.Context, email, username
// 尽力补全:当用户名为空时,使用第三方返回的用户名回填。
if user.Username == "" && username != "" {
user.Username = username
if err := s.userRepo.Update(ctx, user); err != nil {
if err := s.userRepo.Update(ctx, user, UserUpdateFields{Username: true}); err != nil {
logger.LegacyPrintf("service.auth", "[Auth] Failed to update username after oauth login: %v", err)
}
}
@@ -753,7 +753,7 @@ func (s *AuthService) loginOrRegisterOAuthWithTokenPair(ctx context.Context, ema
if user.Username == "" && username != "" {
user.Username = username
if err := s.userRepo.Update(ctx, user); err != nil {
if err := s.userRepo.Update(ctx, user, UserUpdateFields{Username: true}); err != nil {
logger.LegacyPrintf("service.auth", "[Auth] Failed to update username after oauth login: %v", err)
}
}
@@ -1435,7 +1435,9 @@ func (s *AuthService) ResetPassword(ctx context.Context, email, token, newPasswo
user.PasswordHash = hashedPassword
user.TokenVersion++ // Invalidate all existing tokens
if err := s.userRepo.Update(ctx, user); err != nil {
// TokenVersion 无对应数据库列(见 resolvedTokenVersion:由 email+password_hash 指纹推导),
// 写回 password_hash 本身即可让旧 token 失效。
if err := s.userRepo.Update(ctx, user, UserUpdateFields{PasswordHash: true}); err != nil {
logger.LegacyPrintf("service.auth", "[Auth] Database error updating password for user %d: %v", user.ID, err)
return ErrServiceUnavailable
}
@@ -1674,19 +1676,17 @@ func (s *AuthService) RevokeAllUserSessions(ctx context.Context, userID int64) e
}
// RevokeAllUserTokens invalidates both stateless access tokens and refresh sessions.
// Access/refresh token verification both depend on TokenVersion, so bumping it provides
// immediate revocation even if refresh-token cache cleanup later fails.
//
// 注意:users 表没有 token_version 列(resolvedTokenVersion 由 email+password_hash
// 指纹推导),因此对 user.TokenVersion 自增只影响内存副本。之前紧跟其后的整行
// Update 不写任何有效数据,却会用旧快照覆盖并发写入的列,故已移除。
// 会话撤销由下面的 refresh session 清理承担;改密路径通过 password_hash 变化
// 改变指纹,从而使旧 token 失效。
func (s *AuthService) RevokeAllUserTokens(ctx context.Context, userID int64) error {
user, err := s.userRepo.GetByID(ctx, userID)
if err != nil {
if _, err := s.userRepo.GetByID(ctx, userID); err != nil {
return fmt.Errorf("get user: %w", err)
}
user.TokenVersion++
if err := s.userRepo.Update(ctx, user); err != nil {
return fmt.Errorf("update user: %w", err)
}
if err := s.RevokeAllUserSessions(ctx, userID); err != nil {
logger.LegacyPrintf("service.auth", "[Auth] Failed to revoke refresh sessions after token invalidation for user %d: %v", userID, err)
}
@@ -902,7 +902,7 @@ func (s *emailBindUserRepoStub) GetFirstAdmin(context.Context) (*service.User, e
panic("unexpected GetFirstAdmin call")
}
func (s *emailBindUserRepoStub) Update(_ context.Context, user *service.User) error {
func (s *emailBindUserRepoStub) Update(_ context.Context, user *service.User, _ service.UserUpdateFields) error {
s.mu.Lock()
defer s.mu.Unlock()
existing, ok := s.usersByID[user.ID]
@@ -961,6 +961,14 @@ func (s *emailBindUserRepoStub) ExistsByEmail(_ context.Context, email string) (
return ok, nil
}
func (s *emailBindUserRepoStub) AdjustBalance(ctx context.Context, id int64, delta float64) (service.BalanceChange, error) {
panic("unexpected AdjustBalance call")
}
func (s *emailBindUserRepoStub) SetBalance(ctx context.Context, id int64, value float64) (service.BalanceChange, error) {
panic("unexpected SetBalance call")
}
func (s *emailBindUserRepoStub) ExistsByEmailAlias(_ context.Context, email string) (bool, error) {
s.mu.Lock()
defer s.mu.Unlock()
@@ -1297,7 +1297,7 @@ func (s *ContentModerationService) UnbanUser(ctx context.Context, userID int64)
}
if user.Status != StatusActive {
user.Status = StatusActive
if err := s.userRepo.Update(ctx, user); err != nil {
if err := s.userRepo.Update(ctx, user, UserUpdateFields{Status: true}); err != nil {
return nil, fmt.Errorf("update content moderation unban user: %w", err)
}
}
@@ -1832,7 +1832,7 @@ func (s *ContentModerationService) applyFlaggedAccountSideEffects(ctx context.Co
}
if user.Status != StatusDisabled {
user.Status = StatusDisabled
if err := s.userRepo.Update(ctx, user); err != nil {
if err := s.userRepo.Update(ctx, user, UserUpdateFields{Status: true}); err != nil {
slog.Warn("content_moderation.ban_update_user_failed", "user_id", *log.UserID, "error", err)
return false
}
@@ -188,7 +188,7 @@ func (r *contentModerationTestUserRepo) GetFirstAdmin(ctx context.Context) (*Use
panic("unexpected GetFirstAdmin call")
}
func (r *contentModerationTestUserRepo) Update(ctx context.Context, user *User) error {
func (r *contentModerationTestUserRepo) Update(ctx context.Context, user *User, fields UserUpdateFields) error {
if user == nil {
return nil
}
@@ -242,6 +242,14 @@ func (r *contentModerationTestUserRepo) DeductBalance(ctx context.Context, id in
panic("unexpected DeductBalance call")
}
func (r *contentModerationTestUserRepo) AdjustBalance(ctx context.Context, id int64, delta float64) (BalanceChange, error) {
panic("unexpected AdjustBalance call")
}
func (r *contentModerationTestUserRepo) SetBalance(ctx context.Context, id int64, value float64) (BalanceChange, error) {
panic("unexpected SetBalance call")
}
func (r *contentModerationTestUserRepo) UpdateConcurrency(ctx context.Context, id int64, amount int) error {
panic("unexpected UpdateConcurrency call")
}
@@ -154,6 +154,14 @@ func (s *openAIRecordUsageUserRepoStub) DeductBalance(ctx context.Context, id in
return s.deductErr
}
func (s *openAIRecordUsageUserRepoStub) AdjustBalance(ctx context.Context, id int64, delta float64) (BalanceChange, error) {
panic("unexpected AdjustBalance call")
}
func (s *openAIRecordUsageUserRepoStub) SetBalance(ctx context.Context, id int64, value float64) (BalanceChange, error) {
panic("unexpected SetBalance call")
}
type openAIRecordUsageSubRepoStub struct {
UserSubscriptionRepository
+69 -8
View File
@@ -31,6 +31,7 @@ import (
var (
ErrUserNotFound = infraerrors.NotFound("USER_NOT_FOUND", "user not found")
ErrPasswordIncorrect = infraerrors.BadRequest("PASSWORD_INCORRECT", "current password is incorrect")
ErrBalanceNegative = infraerrors.BadRequest("BALANCE_NEGATIVE", "balance cannot be negative")
ErrInsufficientPerms = infraerrors.Forbidden("INSUFFICIENT_PERMISSIONS", "insufficient permissions")
ErrNotifyCodeUserRateLimit = infraerrors.TooManyRequests("NOTIFY_CODE_USER_RATE_LIMIT", "too many verification codes requested, please try again later")
ErrAvatarInvalid = infraerrors.BadRequest("AVATAR_INVALID", "avatar must be a valid image data URL or http(s) URL")
@@ -83,6 +84,48 @@ type UserListFilters struct {
IncludeDeleted bool
}
// UserUpdateFields 声明 UserRepository.Update 允许写回的列。
//
// 未声明的列保持数据库当前值,不会被调用方手里的快照覆盖。用户行上有多条
// 不经过 Update 的原子写入路径(DeductBalance/UpdateBalance 扣加余额、
// UpdateConcurrency、BatchUpdateLimits、UpdateUserLastActiveAt 等),
// status/role 也可能被其他流程并发改写。若 Update 无条件整行回写,
// 一次"读-改-写"就会静默回滚这些并发结果(lost update),
// 因此每个调用方必须显式声明它真正要改的列。
//
// 注意这里没有 balance / total_recharged:余额只能经由 AdjustBalance、
// SetBalance、UpdateBalance、DeductBalance 等原子接口修改,Update 永远不碰它们。
type UserUpdateFields struct {
Email bool
Username bool
Notes bool
PasswordHash bool
Role bool
Status bool
Concurrency bool
RPMLimit bool
SignupSource bool
LastLoginAt bool
LastActiveAt bool
// BalanceNotifySettings 覆盖 balance_notify_enabled / _threshold_type / _threshold。
BalanceNotifySettings bool
// BalanceNotifyExtraEmails 与上一项分开,避免"改通知阈值"覆盖并发的"加通知邮箱"。
BalanceNotifyExtraEmails bool
// AllowedGroups 为 true 时才同步 user_allowed_groups 关联表。
AllowedGroups bool
}
// BalanceChange 记录一次余额变更前后的值。
type BalanceChange struct {
Old float64
New float64
}
// IsEmpty 报告该次 Update 是否不写任何列(此时仓储直接返回,不产生写操作)。
func (f UserUpdateFields) IsEmpty() bool {
return f == UserUpdateFields{}
}
type UserRepository interface {
Create(ctx context.Context, user *User) error
// CreateWithEmailAliasGuard 创建用户,并在邮箱唯一性锁内复查"收件箱身份"是否已被占用
@@ -95,7 +138,8 @@ type UserRepository interface {
GetByIDIncludeDeleted(ctx context.Context, id int64) (*User, error)
GetByEmail(ctx context.Context, email string) (*User, error)
GetFirstAdmin(ctx context.Context) (*User, error)
Update(ctx context.Context, user *User) error
// Update 只写 fields 中显式声明的列,其余列保持库中当前值。
Update(ctx context.Context, user *User, fields UserUpdateFields) error
Delete(ctx context.Context, id int64) error
GetUserAvatar(ctx context.Context, userID int64) (*UserAvatar, error)
UpsertUserAvatar(ctx context.Context, userID int64, input UpsertUserAvatarInput) (*UserAvatar, error)
@@ -109,6 +153,12 @@ type UserRepository interface {
UpdateBalance(ctx context.Context, id int64, amount float64) error
DeductBalance(ctx context.Context, id int64, amount float64) error
// AdjustBalance 原子地把 delta 累加到余额上,并返回变更前后的值。结果为负时
// 拒绝写入并返回 ErrBalanceNegative。管理员的加/扣款必须走这里而不是
// "读余额→算新值→整行写回",否则并发的计费扣款会被旧快照抹掉。
AdjustBalance(ctx context.Context, id int64, delta float64) (BalanceChange, error)
// SetBalance 原子地把余额置为 value(value 必须 >= 0),返回变更前后的值。
SetBalance(ctx context.Context, id int64, value float64) (BalanceChange, error)
UpdateConcurrency(ctx context.Context, id int64, amount int) error
BatchSetConcurrency(ctx context.Context, userIDs []int64, value int) (int, error)
BatchAddConcurrency(ctx context.Context, userIDs []int64, delta int) (int, error)
@@ -443,6 +493,10 @@ func (s *UserService) updateProfile(ctx context.Context, userID int64, req Updat
}
oldConcurrency := user.Concurrency
// fields 只登记本次请求真正带上的字段。余额、状态等列不由这里回写,
// 否则并发的扣费与状态变更会被这份快照回滚。
var fields UserUpdateFields
// 更新字段
if req.Email != nil {
// 检查新邮箱是否已被使用
@@ -454,10 +508,12 @@ func (s *UserService) updateProfile(ctx context.Context, userID int64, req Updat
return nil, oldConcurrency, ErrEmailExists
}
user.Email = *req.Email
fields.Email = true
}
if req.Username != nil {
user.Username = *req.Username
fields.Username = true
}
if req.AvatarURL != nil {
@@ -470,10 +526,12 @@ func (s *UserService) updateProfile(ctx context.Context, userID int64, req Updat
if req.Concurrency != nil {
user.Concurrency = *req.Concurrency
fields.Concurrency = true
}
if req.BalanceNotifyEnabled != nil {
user.BalanceNotifyEnabled = *req.BalanceNotifyEnabled
fields.BalanceNotifySettings = true
}
if req.BalanceNotifyThreshold != nil {
if *req.BalanceNotifyThreshold <= 0 {
@@ -481,9 +539,10 @@ func (s *UserService) updateProfile(ctx context.Context, userID int64, req Updat
} else {
user.BalanceNotifyThreshold = req.BalanceNotifyThreshold
}
fields.BalanceNotifySettings = true
}
if err := s.userRepo.Update(ctx, user); err != nil {
if err := s.userRepo.Update(ctx, user, fields); err != nil {
return nil, oldConcurrency, fmt.Errorf("update user: %w", err)
}
@@ -970,7 +1029,9 @@ func (s *UserService) ChangePassword(ctx context.Context, userID int64, req Chan
// This ensures that any tokens issued before the password change become invalid
user.TokenVersion++
if err := s.userRepo.Update(ctx, user); err != nil {
// TokenVersion 没有对应的数据库列(见 resolvedTokenVersion:它由 email+password_hash
// 指纹推导),改密写回 password_hash 即可让旧 token 失效。
if err := s.userRepo.Update(ctx, user, UserUpdateFields{PasswordHash: true}); err != nil {
return fmt.Errorf("update user: %w", err)
}
@@ -1125,7 +1186,7 @@ func (s *UserService) UpdateStatus(ctx context.Context, userID int64, status str
user.Status = status
if err := s.userRepo.Update(ctx, user); err != nil {
if err := s.userRepo.Update(ctx, user, UserUpdateFields{Status: true}); err != nil {
return fmt.Errorf("update user: %w", err)
}
if s.authCacheInvalidator != nil {
@@ -1285,7 +1346,7 @@ func (s *UserService) addOrVerifyNotifyEmail(ctx context.Context, userID int64,
if strings.EqualFold(e.Email, email) {
if !e.Verified {
user.BalanceNotifyExtraEmails[i].Verified = true
return s.userRepo.Update(ctx, user)
return s.userRepo.Update(ctx, user, UserUpdateFields{BalanceNotifyExtraEmails: true})
}
return nil // Already verified
}
@@ -1298,7 +1359,7 @@ func (s *UserService) addOrVerifyNotifyEmail(ctx context.Context, userID int64,
Disabled: false,
Verified: true,
})
return s.userRepo.Update(ctx, user)
return s.userRepo.Update(ctx, user, UserUpdateFields{BalanceNotifyExtraEmails: true})
}
// RemoveNotifyEmail removes an email from user's extra notification emails.
@@ -1321,7 +1382,7 @@ func (s *UserService) RemoveNotifyEmail(ctx context.Context, userID int64, email
return infraerrors.BadRequest("EMAIL_NOT_FOUND", "notification email not found")
}
user.BalanceNotifyExtraEmails = filtered
return s.userRepo.Update(ctx, user)
return s.userRepo.Update(ctx, user, UserUpdateFields{BalanceNotifyExtraEmails: true})
}
// ToggleNotifyEmail toggles the disabled state of a notification email entry.
@@ -1343,7 +1404,7 @@ func (s *UserService) ToggleNotifyEmail(ctx context.Context, userID int64, email
return infraerrors.BadRequest("EMAIL_NOT_FOUND", "notification email not found")
}
return s.userRepo.Update(ctx, user)
return s.userRepo.Update(ctx, user, UserUpdateFields{BalanceNotifyExtraEmails: true})
}
// notifyVerifyEmailTemplate is the HTML template for notify email verification.
+11 -1
View File
@@ -36,6 +36,7 @@ type mockUserRepo struct {
updateLastActiveAt []time.Time
updateFn func(ctx context.Context, user *User) error
updateCalls int
updateFields []UserUpdateFields
upsertAvatarFn func(ctx context.Context, userID int64, input UpsertUserAvatarInput) (*UserAvatar, error)
upsertAvatarArgs []UpsertUserAvatarInput
deleteAvatarFn func(ctx context.Context, userID int64) error
@@ -108,8 +109,9 @@ func (m *mockUserRepo) GetByID(ctx context.Context, _ int64) (*User, error) {
}
func (m *mockUserRepo) GetByEmail(context.Context, string) (*User, error) { return &User{}, nil }
func (m *mockUserRepo) GetFirstAdmin(context.Context) (*User, error) { return &User{}, nil }
func (m *mockUserRepo) Update(ctx context.Context, user *User) error {
func (m *mockUserRepo) Update(ctx context.Context, user *User, fields UserUpdateFields) error {
m.updateCalls++
m.updateFields = append(m.updateFields, fields)
if m.updateFn != nil {
return m.updateFn(ctx, user)
}
@@ -201,6 +203,14 @@ func (m *mockUserRepo) DeductBalance(ctx context.Context, id int64, amount float
}
return nil
}
func (m *mockUserRepo) AdjustBalance(ctx context.Context, id int64, delta float64) (BalanceChange, error) {
panic("unexpected AdjustBalance call")
}
func (m *mockUserRepo) SetBalance(ctx context.Context, id int64, value float64) (BalanceChange, error) {
panic("unexpected SetBalance call")
}
func (m *mockUserRepo) UpdateConcurrency(context.Context, int64, int) error { return nil }
func (m *mockUserRepo) ExistsByEmail(context.Context, string) (bool, error) { return false, nil }
func (m *mockUserRepo) ExistsByEmailAlias(context.Context, string) (bool, error) {
@@ -0,0 +1,83 @@
//go:build unit
package service
import (
"context"
"testing"
"github.com/stretchr/testify/require"
)
// 这些用例锁死"每个入口只声明自己真正要改的列":
// 任何退回整行回写的改动都会让并发写入被陈旧快照覆盖,并在这里变红。
func TestUpdateProfile_OnlyDeclaresRequestedColumns(t *testing.T) {
username := "renamed"
tests := []struct {
name string
req UpdateProfileRequest
want UserUpdateFields
}{
{
name: "username only",
req: UpdateProfileRequest{Username: &username},
want: UserUpdateFields{Username: true},
},
{
name: "notify settings only",
req: UpdateProfileRequest{BalanceNotifyEnabled: boolPtr(true)},
want: UserUpdateFields{BalanceNotifySettings: true},
},
{
name: "username and notify threshold",
req: UpdateProfileRequest{Username: &username, BalanceNotifyThreshold: float64Ptr(1.5)},
want: UserUpdateFields{Username: true, BalanceNotifySettings: true},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
repo := &mockUserRepo{getByIDUser: &User{ID: 7, Balance: 0.30, Status: StatusActive}}
svc := NewUserService(repo, nil, nil, nil)
_, err := svc.UpdateProfile(context.Background(), 7, tt.req)
require.NoError(t, err)
require.Equal(t, []UserUpdateFields{tt.want}, repo.updateFields)
})
}
}
// 只改头像时用户行没有任何列要写,不应产生一次整行更新。
func TestUpdateProfile_AvatarOnlySkipsUserRowWrite(t *testing.T) {
repo := &mockUserRepo{getByIDUser: &User{ID: 7, Balance: 0.30}}
svc := NewUserService(repo, nil, nil, nil)
avatar := "https://cdn.example.com/a.png"
_, err := svc.UpdateProfile(context.Background(), 7, UpdateProfileRequest{AvatarURL: &avatar})
require.NoError(t, err)
require.Len(t, repo.upsertAvatarArgs, 1, "avatar must still be stored")
require.Equal(t, []UserUpdateFields{{}}, repo.updateFields, "no user column should be declared")
}
func TestChangePassword_OnlyDeclaresPasswordHash(t *testing.T) {
user := &User{ID: 7, Balance: 0.30}
require.NoError(t, user.SetPassword("old-password"))
repo := &mockUserRepo{getByIDUser: user}
svc := NewUserService(repo, nil, nil, nil)
err := svc.ChangePassword(context.Background(), 7, ChangePasswordRequest{
CurrentPassword: "old-password",
NewPassword: "new-password",
})
require.NoError(t, err)
require.Equal(t, []UserUpdateFields{{PasswordHash: true}}, repo.updateFields)
}
func TestUpdateStatus_OnlyDeclaresStatus(t *testing.T) {
repo := &mockUserRepo{getByIDUser: &User{ID: 7, Balance: 0.30, Status: StatusActive}}
svc := NewUserService(repo, nil, nil, nil)
require.NoError(t, svc.UpdateStatus(context.Background(), 7, StatusDisabled))
require.Equal(t, []UserUpdateFields{{Status: true}}, repo.updateFields)
}