diff --git a/common/gin.go b/common/gin.go index 0a7f053..a173588 100644 --- a/common/gin.go +++ b/common/gin.go @@ -170,6 +170,23 @@ func GetContextKeyInt(c *gin.Context, key constant.ContextKey) int { return c.GetInt(string(key)) } +func GetContextKeyInt64(c *gin.Context, key constant.ContextKey) int64 { + value, ok := c.Get(string(key)) + if !ok { + return 0 + } + switch v := value.(type) { + case int64: + return v + case int: + return int64(v) + case float64: + return int64(v) + default: + return 0 + } +} + func GetContextKeyBool(c *gin.Context, key constant.ContextKey) bool { return c.GetBool(string(key)) } diff --git a/controller/billing.go b/controller/billing.go index f40e371..8870f5a 100644 --- a/controller/billing.go +++ b/controller/billing.go @@ -9,8 +9,8 @@ import ( ) func GetSubscription(c *gin.Context) { - var remainQuota int - var usedQuota int + var remainQuota int64 + var usedQuota int64 var err error var token *model.Token var expiredTime int64 @@ -18,8 +18,8 @@ func GetSubscription(c *gin.Context) { tokenId := c.GetInt("token_id") token, err = model.GetTokenById(tokenId) expiredTime = token.ExpiredTime - remainQuota = token.RemainQuota - usedQuota = token.UsedQuota + remainQuota = int64(token.RemainQuota) + usedQuota = int64(token.UsedQuota) } else { userId := c.GetInt("id") remainQuota, err = model.GetUserQuota(userId, false) @@ -69,13 +69,13 @@ func GetSubscription(c *gin.Context) { } func GetUsage(c *gin.Context) { - var quota int + var quota int64 var err error var token *model.Token if common.DisplayTokenStatEnabled { tokenId := c.GetInt("token_id") token, err = model.GetTokenById(tokenId) - quota = token.UsedQuota + quota = int64(token.UsedQuota) } else { userId := c.GetInt("id") quota, err = model.GetUserUsedQuota(userId) diff --git a/controller/user.go b/controller/user.go index c1259c7..b1cdde9 100644 --- a/controller/user.go +++ b/controller/user.go @@ -408,7 +408,7 @@ func GenerateAccessToken(c *gin.Context) { } type TransferAffQuotaRequest struct { - Quota int `json:"quota" binding:"required"` + Quota int64 `json:"quota" binding:"required"` } func TransferAffQuota(c *gin.Context) { @@ -991,19 +991,19 @@ func CreateUser(c *gin.Context) { type ManageRequest struct { Id int `json:"id"` Action string `json:"action"` - Value int `json:"value"` + Value int64 `json:"value"` Mode string `json:"mode"` } -func isValidQuotaOverride(value int) bool { - return value >= 0 && int64(value) <= maxUserQuotaValue +func isValidQuotaOverride(value int64) bool { + return value >= 0 && value <= maxUserQuotaValue } -func isValidQuotaAddition(current int, delta int) bool { - if delta <= 0 || int64(delta) > maxUserQuotaValue { +func isValidQuotaAddition(current int64, delta int64) bool { + if delta <= 0 || delta > maxUserQuotaValue { return false } - return int64(current) <= maxUserQuotaValue-int64(delta) + return current <= maxUserQuotaValue-delta } // ManageUser Only admin user can do this diff --git a/controller/user_setting_test.go b/controller/user_setting_test.go index 156a725..5f2531f 100644 --- a/controller/user_setting_test.go +++ b/controller/user_setting_test.go @@ -217,14 +217,37 @@ func TestRegisterConsumesEmailVerificationCode(t *testing.T) { } func TestQuotaBoundsValidation(t *testing.T) { - legacyInt32Max := 1<<31 - 1 + legacyInt32Max := int64(1<<31 - 1) aboveLegacyInt32Max := legacyInt32Max + 1 + maxQuota := maxUserQuotaValue + overflowingHalfMax := maxQuota/2 + 1 require.True(t, isValidQuotaOverride(0)) require.True(t, isValidQuotaOverride(aboveLegacyInt32Max)) + require.True(t, isValidQuotaOverride(maxQuota)) require.False(t, isValidQuotaOverride(-1)) require.True(t, isValidQuotaAddition(aboveLegacyInt32Max, 1)) require.True(t, isValidQuotaAddition(0, aboveLegacyInt32Max)) + require.True(t, isValidQuotaAddition(maxQuota-1, 1)) + require.True(t, isValidQuotaAddition(0, maxQuota)) + require.False(t, isValidQuotaAddition(maxQuota, 1)) + require.False(t, isValidQuotaAddition(1, maxQuota)) + require.False(t, isValidQuotaAddition(overflowingHalfMax, overflowingHalfMax)) + require.False(t, isValidQuotaAddition(maxQuota, maxQuota)) require.False(t, isValidQuotaAddition(0, 0)) } + +func TestManageUserRejectsQuotaValueAboveInt64Max(t *testing.T) { + recorder := httptest.NewRecorder() + ctx, _ := gin.CreateTestContext(recorder) + ctx.Request = httptest.NewRequest(http.MethodPost, "/api/user/manage", strings.NewReader( + `{"id":1,"action":"add_quota","mode":"override","value":9223372036854775808}`, + )) + ctx.Request.Header.Set("Content-Type", "application/json") + + ManageUser(ctx) + + require.Equal(t, http.StatusOK, recorder.Code) + require.Contains(t, recorder.Body.String(), `"success":false`) +} diff --git a/logger/logger.go b/logger/logger.go index 4a16909..bc1dc8b 100644 --- a/logger/logger.go +++ b/logger/logger.go @@ -119,7 +119,11 @@ func logHelper(ctx context.Context, level string, msg string) { } } -func LogQuota(quota int) string { +type quotaInteger interface { + ~int | ~int64 +} + +func LogQuota[T quotaInteger](quota T) string { // 新逻辑:根据额度展示类型输出 q := float64(quota) switch operation_setting.GetQuotaDisplayType() { @@ -146,7 +150,7 @@ func LogQuota(quota int) string { } } -func FormatQuota(quota int) string { +func FormatQuota[T quotaInteger](quota T) string { q := float64(quota) switch operation_setting.GetQuotaDisplayType() { case operation_setting.QuotaDisplayTypeCNY: diff --git a/model/channel.go b/model/channel.go index 5a14f5e..c0a2933 100644 --- a/model/channel.go +++ b/model/channel.go @@ -854,7 +854,7 @@ func EditChannelByTag(tag string, newTag *string, modelMapping *string, models * func UpdateChannelUsedQuota(id int, quota int) { if common.BatchUpdateEnabled { - addNewRecord(BatchUpdateTypeChannelUsedQuota, id, quota) + addNewRecord(BatchUpdateTypeChannelUsedQuota, id, int64(quota)) return } updateChannelUsedQuota(id, quota) diff --git a/model/main.go b/model/main.go index e2183bf..929895d 100644 --- a/model/main.go +++ b/model/main.go @@ -1,6 +1,7 @@ package model import ( + "database/sql" "errors" "fmt" "log" @@ -666,6 +667,10 @@ func migrateTokenModelLimitsToText() error { return nil } +// migrateUserQuotaColumnsToBigInt runs during startup. On MySQL/PostgreSQL, +// changing large users table columns can hold table-level or rewrite locks for +// a noticeable time, so operators with large deployments should schedule this +// upgrade off-peak or run an equivalent online-DDL migration before booting. func migrateUserQuotaColumnsToBigInt() error { if DB == nil || common.UsingSQLite || !DB.Migrator().HasTable(&User{}) { return nil @@ -679,33 +684,56 @@ func migrateUserQuotaColumnsToBigInt() error { continue } + var backfillSQL string var alterSQL string if common.UsingPostgreSQL { - var dataType string - if err := DB.Raw(`SELECT data_type FROM information_schema.columns + var columnMetadata struct { + DataType string `gorm:"column:data_type"` + IsNullable string `gorm:"column:is_nullable"` + ColumnDefault sql.NullString `gorm:"column:column_default"` + } + if err := DB.Raw(`SELECT data_type AS data_type, is_nullable AS is_nullable, column_default AS column_default + FROM information_schema.columns WHERE table_schema = current_schema() AND table_name = ? AND column_name = ?`, - tableName, columnName).Scan(&dataType).Error; err != nil { + tableName, columnName).Scan(&columnMetadata).Error; err != nil { common.SysLog(fmt.Sprintf("Warning: failed to query metadata for %s.%s: %v", tableName, columnName, err)) - } else if strings.EqualFold(dataType, "bigint") { + } else if strings.EqualFold(strings.TrimSpace(columnMetadata.DataType), "bigint") && + strings.EqualFold(columnMetadata.IsNullable, "NO") && + isZeroColumnDefault(columnMetadata.ColumnDefault) { continue } - alterSQL = fmt.Sprintf(`ALTER TABLE "%s" ALTER COLUMN "%s" TYPE bigint USING "%s"::bigint`, + backfillSQL = fmt.Sprintf(`UPDATE "%s" SET "%s" = 0 WHERE "%s" IS NULL`, tableName, columnName, columnName) + alterSQL = fmt.Sprintf(`ALTER TABLE "%s" ALTER COLUMN "%s" TYPE bigint USING COALESCE("%s", 0)::bigint, ALTER COLUMN "%s" SET NOT NULL, ALTER COLUMN "%s" SET DEFAULT 0`, + tableName, columnName, columnName, columnName, columnName) } else if common.UsingMySQL { - var columnType string - if err := DB.Raw(`SELECT COLUMN_TYPE FROM information_schema.columns + var columnMetadata struct { + ColumnType string `gorm:"column:column_type"` + IsNullable string `gorm:"column:is_nullable"` + ColumnDefault sql.NullString `gorm:"column:column_default"` + } + if err := DB.Raw(`SELECT COLUMN_TYPE AS column_type, IS_NULLABLE AS is_nullable, COLUMN_DEFAULT AS column_default + FROM information_schema.columns WHERE table_schema = DATABASE() AND table_name = ? AND column_name = ?`, - tableName, columnName).Scan(&columnType).Error; err != nil { + tableName, columnName).Scan(&columnMetadata).Error; err != nil { common.SysLog(fmt.Sprintf("Warning: failed to query metadata for %s.%s: %v", tableName, columnName, err)) - } else if strings.HasPrefix(strings.ToLower(columnType), "bigint") { + } else if strings.HasPrefix(strings.ToLower(strings.TrimSpace(columnMetadata.ColumnType)), "bigint") && + strings.EqualFold(columnMetadata.IsNullable, "NO") && + isZeroColumnDefault(columnMetadata.ColumnDefault) { continue } - alterSQL = fmt.Sprintf("ALTER TABLE `%s` MODIFY COLUMN `%s` bigint DEFAULT 0", tableName, columnName) + backfillSQL = fmt.Sprintf("UPDATE `%s` SET `%s` = 0 WHERE `%s` IS NULL", tableName, columnName, columnName) + alterSQL = fmt.Sprintf("ALTER TABLE `%s` MODIFY COLUMN `%s` bigint NOT NULL DEFAULT 0", tableName, columnName) } if alterSQL == "" { continue } + if backfillSQL != "" { + if err := DB.Exec(backfillSQL).Error; err != nil { + return fmt.Errorf("failed to backfill null values for %s.%s: %w", tableName, columnName, err) + } + } if err := DB.Exec(alterSQL).Error; err != nil { return fmt.Errorf("failed to migrate %s.%s to bigint: %w", tableName, columnName, err) } @@ -715,6 +743,20 @@ func migrateUserQuotaColumnsToBigInt() error { return nil } +func isZeroColumnDefault(defaultValue sql.NullString) bool { + if !defaultValue.Valid { + return false + } + + value := strings.ToLower(strings.TrimSpace(defaultValue.String)) + switch value { + case "0", "(0)", "'0'", "0::bigint", "0::integer", "'0'::bigint", "'0'::integer": + return true + default: + return false + } +} + // migrateSubscriptionPlanPriceAmount migrates price_amount column from float/double to decimal(10,6) // This is safe to run multiple times - it checks the column type first func migrateSubscriptionPlanPriceAmount() { diff --git a/model/main_migration_test.go b/model/main_migration_test.go new file mode 100644 index 0000000..a9a441f --- /dev/null +++ b/model/main_migration_test.go @@ -0,0 +1,28 @@ +package model + +import ( + "database/sql" + "testing" +) + +func TestIsZeroColumnDefault(t *testing.T) { + tests := []struct { + name string + value sql.NullString + want bool + }{ + {name: "mysql numeric zero", value: sql.NullString{String: "0", Valid: true}, want: true}, + {name: "postgres bigint zero", value: sql.NullString{String: "0::bigint", Valid: true}, want: true}, + {name: "postgres quoted bigint zero", value: sql.NullString{String: "'0'::bigint", Valid: true}, want: true}, + {name: "missing default", value: sql.NullString{}, want: false}, + {name: "non-zero default", value: sql.NullString{String: "10", Valid: true}, want: false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := isZeroColumnDefault(tt.value); got != tt.want { + t.Fatalf("isZeroColumnDefault(%q) = %v, want %v", tt.value.String, got, tt.want) + } + }) + } +} diff --git a/model/payment_method_guard_test.go b/model/payment_method_guard_test.go index c523331..e4dd23d 100644 --- a/model/payment_method_guard_test.go +++ b/model/payment_method_guard_test.go @@ -17,7 +17,7 @@ func insertUserForPaymentGuardTest(t *testing.T, id int, quota int) { Id: id, Username: "payment_guard_user", Status: common.UserStatusEnabled, - Quota: quota, + Quota: int64(quota), } require.NoError(t, DB.Create(user).Error) } @@ -89,7 +89,7 @@ func countTopUpsForPaymentGuardTest(t *testing.T, tradeNo string) int64 { return count } -func getUserQuotaForPaymentGuardTest(t *testing.T, userID int) int { +func getUserQuotaForPaymentGuardTest(t *testing.T, userID int) int64 { t.Helper() var user User require.NoError(t, DB.Select("quota").Where("id = ?", userID).First(&user).Error) @@ -108,7 +108,7 @@ func TestRechargeWaffoPancake_RejectsMismatchedPaymentMethod(t *testing.T) { topUp := GetTopUpByTradeNo("waffo-pancake-guard") require.NotNil(t, topUp) assert.Equal(t, common.TopUpStatusPending, topUp.Status) - assert.Equal(t, 0, getUserQuotaForPaymentGuardTest(t, 101)) + assert.EqualValues(t, 0, getUserQuotaForPaymentGuardTest(t, 101)) } func TestUpdatePendingTopUpStatus_RejectsMismatchedPaymentProvider(t *testing.T) { @@ -212,7 +212,7 @@ func TestPurchaseSubscriptionWithBalance_InsufficientQuotaDoesNotOverdraw(t *tes err = PurchaseSubscriptionWithBalance(505, plan.Id) require.Error(t, err) - assert.Equal(t, requiredQuota-1, getUserQuotaForPaymentGuardTest(t, 505)) + assert.EqualValues(t, requiredQuota-1, getUserQuotaForPaymentGuardTest(t, 505)) assert.Zero(t, countUserSubscriptionsForPaymentGuardTest(t, 505)) } @@ -235,7 +235,7 @@ func TestRedeem_UsedCodeDoesNotDoubleCredit(t *testing.T) { quota, err = Redeem("redeem-guard-code", 606) require.ErrorIs(t, err, ErrRedeemFailed) assert.Zero(t, quota) - assert.Equal(t, 123, getUserQuotaForPaymentGuardTest(t, 606)) + assert.EqualValues(t, 123, getUserQuotaForPaymentGuardTest(t, 606)) var reloaded Redemption require.NoError(t, DB.Where("id = ?", redemption.Id).First(&reloaded).Error) @@ -282,7 +282,7 @@ func TestRedeemRejectsNonPositiveQuotaWithoutUsingCode(t *testing.T) { quota, err := Redeem("redeem-zero-quota", 609) require.ErrorIs(t, err, ErrRedeemFailed) assert.Zero(t, quota) - assert.Equal(t, 0, getUserQuotaForPaymentGuardTest(t, 609)) + assert.EqualValues(t, 0, getUserQuotaForPaymentGuardTest(t, 609)) var reloaded Redemption require.NoError(t, DB.Where("id = ?", redemption.Id).First(&reloaded).Error) @@ -309,7 +309,7 @@ func TestRechargeCreemRejectsZeroQuotaBeforeCompletingOrder(t *testing.T) { err := RechargeCreem("creem-zero-quota", "", "", "127.0.0.1") require.Error(t, err) assert.Equal(t, common.TopUpStatusPending, getTopUpStatusForPaymentGuardTest(t, "creem-zero-quota")) - assert.Equal(t, 0, getUserQuotaForPaymentGuardTest(t, 610)) + assert.EqualValues(t, 0, getUserQuotaForPaymentGuardTest(t, 610)) } func TestRechargeCreemSkipsDuplicateCustomerEmailBinding(t *testing.T) { @@ -337,7 +337,7 @@ func TestRechargeCreemSkipsDuplicateCustomerEmailBinding(t *testing.T) { require.NoError(t, DB.First(&got, 612).Error) assert.Empty(t, got.Email) assert.Empty(t, got.NormalizedEmail) - assert.Equal(t, 2, got.Quota) + assert.EqualValues(t, 2, got.Quota) assert.Equal(t, common.TopUpStatusSuccess, getTopUpStatusForPaymentGuardTest(t, "creem-duplicate-email")) count, err := CountUsersByEmail("taken@example.com") diff --git a/model/subscription.go b/model/subscription.go index 910f118..c2ea5b6 100644 --- a/model/subscription.go +++ b/model/subscription.go @@ -789,7 +789,7 @@ func PurchaseSubscriptionWithBalance(userId int, planId int) error { if err := withRowLock(tx).Where("id = ?", userId).First(&user).Error; err != nil { return err } - if requiredQuota > 0 && user.Quota < requiredQuota { + if requiredQuota > 0 && user.Quota < int64(requiredQuota) { return errors.New("余额不足") } if requiredQuota > 0 { diff --git a/model/token.go b/model/token.go index c5f89e5..7f7b249 100644 --- a/model/token.go +++ b/model/token.go @@ -413,7 +413,7 @@ func IncreaseTokenQuota(tokenId int, key string, quota int) (err error) { }) } if common.BatchUpdateEnabled { - addNewRecord(BatchUpdateTypeTokenQuota, tokenId, quota) + addNewRecord(BatchUpdateTypeTokenQuota, tokenId, int64(quota)) return nil } return increaseTokenQuota(tokenId, quota) @@ -443,7 +443,7 @@ func DecreaseTokenQuota(id int, key string, quota int) (err error) { }) } if common.BatchUpdateEnabled { - addNewRecord(BatchUpdateTypeTokenQuota, id, -quota) + addNewRecord(BatchUpdateTypeTokenQuota, id, int64(-quota)) return nil } return decreaseTokenQuota(id, quota) diff --git a/model/user.go b/model/user.go index 8775229..c3b7097 100644 --- a/model/user.go +++ b/model/user.go @@ -64,14 +64,14 @@ type User struct { TelegramId string `json:"telegram_id" gorm:"column:telegram_id;index"` VerificationCode string `json:"verification_code" gorm:"-:all"` // this field is only for Email verification, don't save it to database! AccessToken *string `json:"-" gorm:"type:char(32);column:access_token;uniqueIndex"` // this token is for system management - Quota int `json:"quota" gorm:"type:bigint;default:0"` - UsedQuota int `json:"used_quota" gorm:"type:bigint;default:0;column:used_quota"` // used quota + Quota int64 `json:"quota" gorm:"type:bigint;default:0"` + UsedQuota int64 `json:"used_quota" gorm:"type:bigint;default:0;column:used_quota"` // used quota RequestCount int `json:"request_count" gorm:"type:int;default:0;"` // request number Group string `json:"group" gorm:"type:varchar(64);default:'default'"` AffCode string `json:"aff_code" gorm:"type:varchar(32);column:aff_code;uniqueIndex"` AffCount int `json:"aff_count" gorm:"type:int;default:0;column:aff_count"` - AffQuota int `json:"aff_quota" gorm:"type:bigint;default:0;column:aff_quota"` // 邀请剩余额度 - AffHistoryQuota int `json:"aff_history_quota" gorm:"type:bigint;default:0;column:aff_history"` // 邀请历史额度 + AffQuota int64 `json:"aff_quota" gorm:"type:bigint;default:0;column:aff_quota"` // 邀请剩余额度 + AffHistoryQuota int64 `json:"aff_history_quota" gorm:"type:bigint;default:0;column:aff_history"` // 邀请历史额度 InviterId int `json:"inviter_id" gorm:"type:int;column:inviter_id;index"` DeletedAt gorm.DeletedAt `gorm:"index"` LinuxDOId string `json:"linux_do_id" gorm:"column:linux_do_id;index"` @@ -584,12 +584,12 @@ func inviteUser(inviterId int) (err error) { return err } user.AffCount++ - user.AffQuota += common.QuotaForInviter - user.AffHistoryQuota += common.QuotaForInviter + user.AffQuota += int64(common.QuotaForInviter) + user.AffHistoryQuota += int64(common.QuotaForInviter) return DB.Save(user).Error } -func (user *User) TransferAffQuotaToQuota(quota int) error { +func (user *User) TransferAffQuotaToQuota(quota int64) error { // 检查quota是否小于最小额度 if float64(quota) < common.QuotaPerUnit { return fmt.Errorf("转移额度最小为%s!", logger.LogQuota(int(common.QuotaPerUnit))) @@ -692,7 +692,7 @@ func (user *User) Insert(inviterId int) error { if err := user.prepareForInsert(tx); err != nil { return err } - user.Quota = common.QuotaForNewUser + user.Quota = int64(common.QuotaForNewUser) user.AffCode = common.GetRandomString(4) // 初始化用户设置,包括默认的边栏配置 @@ -749,7 +749,7 @@ func (user *User) InsertWithTx(tx *gorm.DB, inviterId int) error { if err := user.prepareForInsert(tx); err != nil { return err } - user.Quota = common.QuotaForNewUser + user.Quota = int64(common.QuotaForNewUser) user.AffCode = common.GetRandomString(4) // 初始化用户设置 @@ -1336,7 +1336,7 @@ func ValidateAccessToken(token string) (*User, error) { } // 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 int64, err error) { defer func() { // Update Redis cache asynchronously on successful DB read if shouldUpdateRedis(fromDB, err) { @@ -1363,7 +1363,7 @@ func GetUserQuota(id int, fromDB bool) (quota int, err error) { return quota, nil } -func GetUserUsedQuota(id int) (quota int, err error) { +func GetUserUsedQuota(id int) (quota int64, err error) { err = DB.Model(&User{}).Where("id = ?", id).Select("used_quota").Find("a).Error return quota, err } @@ -1439,24 +1439,29 @@ func GetUserSetting(id int, fromDB bool) (settingMap dto.UserSetting, err error) return userBase.GetSetting(), nil } -func IncreaseUserQuota(id int, quota int, db bool) (err error) { - if quota < 0 { +type quotaDeltaInteger interface { + ~int | ~int64 +} + +func IncreaseUserQuota[T quotaDeltaInteger](id int, quota T, db bool) (err error) { + delta := int64(quota) + if delta < 0 { return errors.New("quota 不能为负数!") } gopool.Go(func() { - err := cacheIncrUserQuota(id, int64(quota)) + err := cacheIncrUserQuota(id, delta) if err != nil { common.SysLog("failed to increase user quota: " + err.Error()) } }) if !db && common.BatchUpdateEnabled { - addNewRecord(BatchUpdateTypeUserQuota, id, quota) + addNewRecord(BatchUpdateTypeUserQuota, id, delta) return nil } - return increaseUserQuota(id, quota) + return increaseUserQuota(id, delta) } -func increaseUserQuota(id int, quota int) (err error) { +func increaseUserQuota(id int, quota int64) (err error) { err = DB.Model(&User{}).Where("id = ?", id).Update("quota", gorm.Expr("quota + ?", quota)).Error if err != nil { return err @@ -1464,24 +1469,25 @@ func increaseUserQuota(id int, quota int) (err error) { return err } -func DecreaseUserQuota(id int, quota int, db bool) (err error) { - if quota < 0 { +func DecreaseUserQuota[T quotaDeltaInteger](id int, quota T, db bool) (err error) { + delta := int64(quota) + if delta < 0 { return errors.New("quota 不能为负数!") } gopool.Go(func() { - err := cacheDecrUserQuota(id, int64(quota)) + err := cacheDecrUserQuota(id, delta) if err != nil { common.SysLog("failed to decrease user quota: " + err.Error()) } }) if !db && common.BatchUpdateEnabled { - addNewRecord(BatchUpdateTypeUserQuota, id, -quota) + addNewRecord(BatchUpdateTypeUserQuota, id, -delta) return nil } - return decreaseUserQuota(id, quota) + return decreaseUserQuota(id, delta) } -func decreaseUserQuota(id int, quota int) (err error) { +func decreaseUserQuota(id int, quota int64) (err error) { err = DB.Model(&User{}).Where("id = ?", id).Update("quota", gorm.Expr("quota - ?", quota)).Error if err != nil { return err @@ -1489,7 +1495,7 @@ func decreaseUserQuota(id int, quota int) (err error) { return err } -func DeltaUpdateUserQuota(id int, delta int) (err error) { +func DeltaUpdateUserQuota[T quotaDeltaInteger](id int, delta T) (err error) { if delta == 0 { return nil } @@ -1518,7 +1524,7 @@ func UpdateUserLastLoginAt(id int) { func UpdateUserUsedQuotaAndRequestCount(id int, quota int) { if common.BatchUpdateEnabled { - addNewRecord(BatchUpdateTypeUsedQuota, id, quota) + addNewRecord(BatchUpdateTypeUsedQuota, id, int64(quota)) addNewRecord(BatchUpdateTypeRequestCount, id, 1) return } @@ -1543,7 +1549,7 @@ func updateUserUsedQuotaAndRequestCount(id int, quota int, count int) { //} } -func updateUserQuotaUsedQuotaAndRequestCount(id int, quota int, usedQuota int, requestCount int) { +func updateUserQuotaUsedQuotaAndRequestCount(id int, quota int64, usedQuota int64, requestCount int64) { if quota == 0 && usedQuota == 0 && requestCount == 0 { return } diff --git a/model/user_cache.go b/model/user_cache.go index 867f040..9689a08 100644 --- a/model/user_cache.go +++ b/model/user_cache.go @@ -18,7 +18,7 @@ type UserBase struct { Id int `json:"id"` Group string `json:"group"` Email string `json:"email"` - Quota int `json:"quota"` + Quota int64 `json:"quota"` Role int `json:"role"` Status int `json:"status"` Username string `json:"username"` @@ -184,7 +184,7 @@ func getUserGroupCache(userId int) (string, error) { return cache.Group, nil } -func getUserQuotaCache(userId int) (int, error) { +func getUserQuotaCache(userId int) (int64, error) { cache, err := GetUserCache(userId) if err != nil { return 0, err @@ -228,7 +228,7 @@ func updateUserStatusCache(userId int, status bool) error { return common.RedisHSetField(getUserCacheKey(userId), "Status", fmt.Sprintf("%d", statusInt)) } -func updateUserQuotaCache(userId int, quota int) error { +func updateUserQuotaCache(userId int, quota int64) error { if !common.RedisEnabled { return nil } diff --git a/model/user_update_test.go b/model/user_update_test.go index a0fab97..7ad3b39 100644 --- a/model/user_update_test.go +++ b/model/user_update_test.go @@ -190,8 +190,8 @@ func TestUserUpdateDoesNotOverwriteAccountingFields(t *testing.T) { var got User require.NoError(t, DB.First(&got, user.Id).Error) assert.Equal(t, "after", got.DisplayName) - assert.Equal(t, 600, got.Quota) - assert.Equal(t, 420, got.UsedQuota) + assert.EqualValues(t, 600, got.Quota) + assert.EqualValues(t, 420, got.UsedQuota) assert.Equal(t, 4, got.RequestCount) } @@ -225,8 +225,8 @@ func TestUserUpdatePersistsZeroValueProfileFields(t *testing.T) { require.NoError(t, DB.First(&got, user.Id).Error) assert.Empty(t, got.DisplayName) assert.Zero(t, got.AffCount) - assert.Equal(t, 1000, got.Quota) - assert.Equal(t, 20, got.UsedQuota) + assert.EqualValues(t, 1000, got.Quota) + assert.EqualValues(t, 20, got.UsedQuota) assert.Equal(t, 3, got.RequestCount) } @@ -303,8 +303,8 @@ func TestUpdateUserSettingOnlyUpdatesSetting(t *testing.T) { var got User require.NoError(t, DB.First(&got, user.Id).Error) - assert.Equal(t, 750, got.Quota) - assert.Equal(t, 270, got.UsedQuota) + assert.EqualValues(t, 750, got.Quota) + assert.EqualValues(t, 270, got.UsedQuota) assert.Equal(t, 4, got.RequestCount) assert.Equal(t, "zh", got.GetSetting().Language) @@ -340,8 +340,8 @@ func TestUpdateUserSettingOnlyUpdatesSettingMySQL(t *testing.T) { var got User require.NoError(t, DB.First(&got, user.Id).Error) - assert.Equal(t, 750, got.Quota) - assert.Equal(t, 270, got.UsedQuota) + assert.EqualValues(t, 750, got.Quota) + assert.EqualValues(t, 270, got.UsedQuota) assert.Equal(t, 4, got.RequestCount) assert.Equal(t, "zh", got.GetSetting().Language) diff --git a/model/utils.go b/model/utils.go index ce32ae5..7f93198 100644 --- a/model/utils.go +++ b/model/utils.go @@ -20,12 +20,12 @@ const ( BatchUpdateTypeCount // if you add a new type, you need to add a new map and a new lock ) -var batchUpdateStores []map[int]int +var batchUpdateStores []map[int]int64 var batchUpdateLocks []sync.Mutex func init() { for i := 0; i < BatchUpdateTypeCount; i++ { - batchUpdateStores = append(batchUpdateStores, make(map[int]int)) + batchUpdateStores = append(batchUpdateStores, make(map[int]int64)) batchUpdateLocks = append(batchUpdateLocks, sync.Mutex{}) } } @@ -39,7 +39,7 @@ func InitBatchUpdater() { }) } -func addNewRecord(type_ int, id int, value int) { +func addNewRecord(type_ int, id int, value int64) { batchUpdateLocks[type_].Lock() defer batchUpdateLocks[type_].Unlock() if _, ok := batchUpdateStores[type_][id]; !ok { @@ -67,11 +67,11 @@ func batchUpdate() { } common.SysLog("batch update started") - stores := make([]map[int]int, BatchUpdateTypeCount) + stores := make([]map[int]int64, BatchUpdateTypeCount) for i := 0; i < BatchUpdateTypeCount; i++ { batchUpdateLocks[i].Lock() stores[i] = batchUpdateStores[i] - batchUpdateStores[i] = make(map[int]int) + batchUpdateStores[i] = make(map[int]int64) batchUpdateLocks[i].Unlock() } @@ -82,12 +82,12 @@ func batchUpdate() { for key, value := range store { switch i { case BatchUpdateTypeTokenQuota: - err := increaseTokenQuota(key, value) + err := increaseTokenQuota(key, int(value)) if err != nil { common.SysLog("failed to batch update token quota: " + err.Error()) } case BatchUpdateTypeChannelUsedQuota: - updateChannelUsedQuota(key, value) + updateChannelUsedQuota(key, int(value)) } } } diff --git a/relay/common/relay_info.go b/relay/common/relay_info.go index ff4aa7c..f2aaeea 100644 --- a/relay/common/relay_info.go +++ b/relay/common/relay_info.go @@ -117,7 +117,7 @@ type RelayInfo struct { ReasoningEffort string UserSetting dto.UserSetting UserEmail string - UserQuota int + UserQuota int64 RelayFormat types.RelayFormat SendResponseCount int ReceivedResponseCount int @@ -471,7 +471,7 @@ func genBaseRelayInfo(c *gin.Context, request dto.Request) *RelayInfo { UserId: common.GetContextKeyInt(c, constant.ContextKeyUserId), UsingGroup: common.GetContextKeyString(c, constant.ContextKeyUsingGroup), UserGroup: common.GetContextKeyString(c, constant.ContextKeyUserGroup), - UserQuota: common.GetContextKeyInt(c, constant.ContextKeyUserQuota), + UserQuota: common.GetContextKeyInt64(c, constant.ContextKeyUserQuota), UserEmail: common.GetContextKeyString(c, constant.ContextKeyUserEmail), OriginModelName: common.GetContextKeyString(c, constant.ContextKeyOriginalModel), diff --git a/relay/mjproxy_handler.go b/relay/mjproxy_handler.go index 12c9288..2d2504e 100644 --- a/relay/mjproxy_handler.go +++ b/relay/mjproxy_handler.go @@ -210,7 +210,7 @@ func RelaySwapFace(c *gin.Context, info *relaycommon.RelayInfo) *dto.MidjourneyR } } - if userQuota-priceData.Quota < 0 { + if userQuota-int64(priceData.Quota) < 0 { return &dto.MidjourneyResponse{ Code: 4, Description: "quota_not_enough", @@ -517,7 +517,7 @@ func RelayMidjourneySubmit(c *gin.Context, relayInfo *relaycommon.RelayInfo) *dt } } - if consumeQuota && userQuota-priceData.Quota < 0 { + if consumeQuota && userQuota-int64(priceData.Quota) < 0 { return &dto.MidjourneyResponse{ Code: 4, Description: "quota_not_enough", diff --git a/service/billing_session.go b/service/billing_session.go index 2d2eaf6..409cc39 100644 --- a/service/billing_session.go +++ b/service/billing_session.go @@ -305,7 +305,7 @@ func (s *BillingSession) shouldTrust(c *gin.Context) bool { switch s.funding.Source() { case BillingSourceWallet: - return s.relayInfo.UserQuota > trustQuota + return s.relayInfo.UserQuota > int64(trustQuota) case BillingSourceSubscription: // 订阅不能启用信任旁路。原因: // 1. PreConsumeUserSubscription 要求 amount>0 来创建预扣记录并锁定订阅 @@ -361,7 +361,7 @@ func NewBillingSession(c *gin.Context, relayInfo *relaycommon.RelayInfo, preCons types.ErrorCodeInsufficientUserQuota, http.StatusForbidden, types.ErrOptionWithSkipRetry(), types.ErrOptionWithNoRecordErrorLog()) } - if userQuota-preConsumedQuota < 0 { + if userQuota-int64(preConsumedQuota) < 0 { return nil, types.NewErrorWithStatusCode( fmt.Errorf("预扣费额度失败, 用户剩余额度: %s, 需要预扣费额度: %s", logger.FormatQuota(userQuota), logger.FormatQuota(preConsumedQuota)), types.ErrorCodeInsufficientUserQuota, http.StatusForbidden, diff --git a/service/pre_consume_quota.go b/service/pre_consume_quota.go index ebbaeeb..7928da0 100644 --- a/service/pre_consume_quota.go +++ b/service/pre_consume_quota.go @@ -38,14 +38,14 @@ func PreConsumeQuota(c *gin.Context, preConsumedQuota int, relayInfo *relaycommo if userQuota <= 0 { return types.NewErrorWithStatusCode(fmt.Errorf("用户额度不足, 剩余额度: %s", logger.FormatQuota(userQuota)), types.ErrorCodeInsufficientUserQuota, http.StatusForbidden, types.ErrOptionWithSkipRetry(), types.ErrOptionWithNoRecordErrorLog()) } - if userQuota-preConsumedQuota < 0 { + if userQuota-int64(preConsumedQuota) < 0 { return types.NewErrorWithStatusCode(fmt.Errorf("预扣费额度失败, 用户剩余额度: %s, 需要预扣费额度: %s", logger.FormatQuota(userQuota), logger.FormatQuota(preConsumedQuota)), types.ErrorCodeInsufficientUserQuota, http.StatusForbidden, types.ErrOptionWithSkipRetry(), types.ErrOptionWithNoRecordErrorLog()) } trustQuota := common.GetTrustQuota() relayInfo.UserQuota = userQuota - if userQuota > trustQuota { + if userQuota > int64(trustQuota) { // 用户额度充足,判断令牌额度是否充足 if !relayInfo.TokenUnlimited { // 非无限令牌,判断令牌额度是否充足 @@ -72,7 +72,7 @@ func PreConsumeQuota(c *gin.Context, preConsumedQuota int, relayInfo *relaycommo if err != nil { return types.NewError(err, types.ErrorCodeUpdateDataError, types.ErrOptionWithSkipRetry()) } - logger.LogInfo(c, fmt.Sprintf("用户 %d 预扣费 %s, 预扣费后剩余额度: %s", relayInfo.UserId, logger.FormatQuota(preConsumedQuota), logger.FormatQuota(userQuota-preConsumedQuota))) + logger.LogInfo(c, fmt.Sprintf("用户 %d 预扣费 %s, 预扣费后剩余额度: %s", relayInfo.UserId, logger.FormatQuota(preConsumedQuota), logger.FormatQuota(userQuota-int64(preConsumedQuota)))) } relayInfo.FinalPreConsumedQuota = preConsumedQuota return nil diff --git a/service/quota.go b/service/quota.go index bdca2bb..8b1412f 100644 --- a/service/quota.go +++ b/service/quota.go @@ -138,7 +138,7 @@ func PreWssConsumeQuota(ctx *gin.Context, relayInfo *relaycommon.RelayInfo, usag quota := calculateAudioQuota(quotaInfo) - if userQuota < quota { + if userQuota < int64(quota) { return fmt.Errorf("user quota is not enough, user quota: %s, need quota: %s", logger.FormatQuota(userQuota), logger.FormatQuota(quota)) } @@ -477,7 +477,7 @@ func checkAndSendQuotaNotify(relayInfo *relaycommon.RelayInfo, quota int, preCon //noMoreQuota := userCache.Quota-(quota+preConsumedQuota) <= 0 quotaTooLow := false consumeQuota := quota + preConsumedQuota - if relayInfo.UserQuota-consumeQuota < threshold { + if relayInfo.UserQuota-int64(consumeQuota) < int64(threshold) { quotaTooLow = true } if quotaTooLow { diff --git a/service/task_billing_test.go b/service/task_billing_test.go index 729f2ae..c7ae51a 100644 --- a/service/task_billing_test.go +++ b/service/task_billing_test.go @@ -85,7 +85,7 @@ func truncate(t *testing.T) { func seedUser(t *testing.T, id int, quota int) { t.Helper() - user := &model.User{Id: id, Username: "test_user", Quota: quota, Status: common.UserStatusEnabled} + user := &model.User{Id: id, Username: "test_user", Quota: int64(quota), Status: common.UserStatusEnabled} require.NoError(t, model.DB.Create(user).Error) } @@ -154,7 +154,7 @@ func makeTask(userId, channelId, quota, tokenId int, billingSource string, subsc // Read-back helpers // --------------------------------------------------------------------------- -func getUserQuota(t *testing.T, id int) int { +func getUserQuota(t *testing.T, id int) int64 { t.Helper() var user model.User require.NoError(t, model.DB.Select("quota").Where("id = ?", id).First(&user).Error) @@ -227,7 +227,7 @@ func TestRefundTaskQuota_Wallet(t *testing.T) { RefundTaskQuota(ctx, task, "task failed: upstream error") // User quota should increase by preConsumed - assert.Equal(t, initQuota+preConsumed, getUserQuota(t, userID)) + assert.EqualValues(t, initQuota+preConsumed, getUserQuota(t, userID)) // Token remain_quota should increase, used_quota should decrease assert.Equal(t, tokenRemain+preConsumed, getTokenRemainQuota(t, tokenID)) @@ -282,7 +282,7 @@ func TestRefundTaskQuota_ZeroQuota(t *testing.T) { RefundTaskQuota(ctx, task, "zero quota task") // No change to user quota - assert.Equal(t, 5000, getUserQuota(t, userID)) + assert.EqualValues(t, 5000, getUserQuota(t, userID)) // No log created assert.Equal(t, int64(0), countLogs(t)) @@ -303,7 +303,7 @@ func TestRefundTaskQuota_NoToken(t *testing.T) { RefundTaskQuota(ctx, task, "no token task failed") // User quota refunded - assert.Equal(t, initQuota+preConsumed, getUserQuota(t, userID)) + assert.EqualValues(t, initQuota+preConsumed, getUserQuota(t, userID)) // Log created log := getLastLog(t) @@ -333,7 +333,7 @@ func TestRecalculate_PositiveDelta(t *testing.T) { RecalculateTaskQuota(ctx, task, actualQuota, "adaptor adjustment") // User quota should decrease by the delta (1000 additional charge) - assert.Equal(t, initQuota-(actualQuota-preConsumed), getUserQuota(t, userID)) + assert.EqualValues(t, initQuota-(actualQuota-preConsumed), getUserQuota(t, userID)) // Token should also be charged the delta assert.Equal(t, tokenRemain-(actualQuota-preConsumed), getTokenRemainQuota(t, tokenID)) @@ -412,7 +412,7 @@ func TestRecalculate_NegativeDelta(t *testing.T) { RecalculateTaskQuota(ctx, task, actualQuota, "adaptor adjustment") // User quota should increase by abs(delta) = 2000 (refund overpayment) - assert.Equal(t, initQuota+(preConsumed-actualQuota), getUserQuota(t, userID)) + assert.EqualValues(t, initQuota+(preConsumed-actualQuota), getUserQuota(t, userID)) // Token should be refunded the difference assert.Equal(t, tokenRemain+(preConsumed-actualQuota), getTokenRemainQuota(t, tokenID)) @@ -441,7 +441,7 @@ func TestRecalculate_ZeroDelta(t *testing.T) { RecalculateTaskQuota(ctx, task, preConsumed, "exact match") // No change to user quota - assert.Equal(t, initQuota, getUserQuota(t, userID)) + assert.EqualValues(t, initQuota, getUserQuota(t, userID)) // No log created (delta is zero) assert.Equal(t, int64(0), countLogs(t)) @@ -461,7 +461,7 @@ func TestRecalculate_ActualQuotaZero(t *testing.T) { RecalculateTaskQuota(ctx, task, 0, "zero actual") // No change (early return) - assert.Equal(t, initQuota, getUserQuota(t, userID)) + assert.EqualValues(t, initQuota, getUserQuota(t, userID)) assert.Equal(t, int64(0), countLogs(t)) } @@ -575,7 +575,7 @@ func TestCASGuardedRefund_Win(t *testing.T) { assert.EqualValues(t, model.TaskStatusFailure, reloaded.Status) // Refund should have happened - assert.Equal(t, initQuota+preConsumed, getUserQuota(t, userID)) + assert.EqualValues(t, initQuota+preConsumed, getUserQuota(t, userID)) assert.Equal(t, tokenRemain+preConsumed, getTokenRemainQuota(t, tokenID)) log := getLastLog(t) @@ -608,7 +608,7 @@ func TestCASGuardedRefund_Lose(t *testing.T) { simulatePollBilling(ctx, task, model.TaskStatus(model.TaskStatusFailure), 0) // CAS lost: user quota should NOT change (no double refund) - assert.Equal(t, initQuota, getUserQuota(t, userID)) + assert.EqualValues(t, initQuota, getUserQuota(t, userID)) assert.Equal(t, tokenRemain, getTokenRemainQuota(t, tokenID)) // No billing log should be created @@ -640,7 +640,7 @@ func TestCASGuardedSettle_Win(t *testing.T) { assert.EqualValues(t, model.TaskStatusSuccess, reloaded.Status) // Settlement should refund the over-charge (5000 - 3000 = 2000 back to user) - assert.Equal(t, initQuota+(preConsumed-actualQuota), getUserQuota(t, userID)) + assert.EqualValues(t, initQuota+(preConsumed-actualQuota), getUserQuota(t, userID)) assert.Equal(t, tokenRemain+(preConsumed-actualQuota), getTokenRemainQuota(t, tokenID)) // task.Quota should be updated to actualQuota @@ -666,7 +666,7 @@ func TestNonTerminalUpdate_NoBilling(t *testing.T) { simulatePollBilling(ctx, task, model.TaskStatus(model.TaskStatusInProgress), 0) // User quota should NOT change - assert.Equal(t, initQuota, getUserQuota(t, userID)) + assert.EqualValues(t, initQuota, getUserQuota(t, userID)) // No billing log assert.Equal(t, int64(0), countLogs(t)) @@ -719,7 +719,7 @@ func TestSettle_PerCallBilling_SkipsAdaptorAdjust(t *testing.T) { settleTaskBillingOnComplete(ctx, adaptor, task, taskResult) // Per-call: no adjustment despite adaptor returning 2000 - assert.Equal(t, initQuota, getUserQuota(t, userID)) + assert.EqualValues(t, initQuota, getUserQuota(t, userID)) assert.Equal(t, tokenRemain, getTokenRemainQuota(t, tokenID)) assert.Equal(t, preConsumed, task.Quota) assert.Equal(t, int64(0), countLogs(t)) @@ -746,7 +746,7 @@ func TestSettle_PerCallBilling_SkipsTotalTokens(t *testing.T) { settleTaskBillingOnComplete(ctx, adaptor, task, taskResult) // Per-call: no recalculation by tokens - assert.Equal(t, initQuota, getUserQuota(t, userID)) + assert.EqualValues(t, initQuota, getUserQuota(t, userID)) assert.Equal(t, tokenRemain, getTokenRemainQuota(t, tokenID)) assert.Equal(t, preConsumed, task.Quota) assert.Equal(t, int64(0), countLogs(t)) @@ -774,7 +774,7 @@ func TestSettle_NonPerCall_AdaptorAdjustWorks(t *testing.T) { settleTaskBillingOnComplete(ctx, adaptor, task, taskResult) // Non-per-call: adaptor adjustment applies (refund 2000) - assert.Equal(t, initQuota+(preConsumed-adaptorQuota), getUserQuota(t, userID)) + assert.EqualValues(t, initQuota+(preConsumed-adaptorQuota), getUserQuota(t, userID)) assert.Equal(t, tokenRemain+(preConsumed-adaptorQuota), getTokenRemainQuota(t, tokenID)) assert.Equal(t, adaptorQuota, task.Quota) diff --git a/service/text_quota_test.go b/service/text_quota_test.go index 72e8266..c0305af 100644 --- a/service/text_quota_test.go +++ b/service/text_quota_test.go @@ -172,7 +172,7 @@ func TestPostTextConsumeQuotaUpdatesUsageStatsForStreamFallback(t *testing.T) { var user model.User require.NoError(t, model.DB.Select("used_quota", "request_count").Where("id = ?", userID).First(&user).Error) - require.Equal(t, fallbackQuota, user.UsedQuota) + require.EqualValues(t, fallbackQuota, user.UsedQuota) require.Equal(t, 1, user.RequestCount) var channel model.Channel @@ -225,7 +225,7 @@ func TestPostTextConsumeQuotaUsesFallbackForNilUsageWithEstimate(t *testing.T) { var user model.User require.NoError(t, model.DB.Select("used_quota", "request_count").Where("id = ?", userID).First(&user).Error) - require.Equal(t, fallbackQuota, user.UsedQuota) + require.EqualValues(t, fallbackQuota, user.UsedQuota) require.Equal(t, 1, user.RequestCount) var channel model.Channel