diff --git a/model/rebate_record.go b/model/rebate_record.go index 97902ff9..773cf1ce 100644 --- a/model/rebate_record.go +++ b/model/rebate_record.go @@ -4,10 +4,10 @@ import "time" type RebateRecord struct { Id int `json:"id" gorm:"primaryKey;autoIncrement"` - InviterId int `json:"inviter_id" gorm:"index;not null"` - InviteeId int `json:"invitee_id" gorm:"index;not null"` - OrderId int `json:"order_id" gorm:"index"` - OrderType string `json:"order_type" gorm:"type:varchar(16);not null"` + InviterId int `json:"inviter_id" gorm:"index;uniqueIndex:idx_rebate_order_invite,priority:3;not null"` + InviteeId int `json:"invitee_id" gorm:"index;uniqueIndex:idx_rebate_order_invite,priority:4;not null"` + OrderId int `json:"order_id" gorm:"index;uniqueIndex:idx_rebate_order_invite,priority:2"` + OrderType string `json:"order_type" gorm:"type:varchar(16);uniqueIndex:idx_rebate_order_invite,priority:1;not null"` OrderAmount int `json:"order_amount" gorm:"not null"` RebateAmount int `json:"rebate_amount" gorm:"not null"` RatePercent int `json:"rate_percent" gorm:"not null"` diff --git a/service/rebate.go b/service/rebate.go new file mode 100644 index 00000000..d8d72529 --- /dev/null +++ b/service/rebate.go @@ -0,0 +1,151 @@ +package service + +import ( + "errors" + "strconv" + "time" + + "github.com/QuantumNous/new-api/common" + "github.com/QuantumNous/new-api/model" + "gorm.io/gorm" + "gorm.io/gorm/clause" +) + +func IsRebateEnabled() bool { + enabled, err := strconv.ParseBool(readRebateOption(model.AffRebateEnabledKey)) + return err == nil && enabled +} + +func GetGlobalRebateRatePercent() int { + return parseNonNegativeRebateInt(readRebateOption(model.AffRebateRatePercentKey)) +} + +func GetGlobalRebateFrozenDays() int { + return parseNonNegativeRebateInt(readRebateOption(model.AffRebateFrozenDaysKey)) +} + +func GetEffectiveRebateRate(inviter *model.User) int { + if inviter != nil && inviter.AffRebateRatePercent > 0 { + return inviter.AffRebateRatePercent + } + return GetGlobalRebateRatePercent() +} + +func GetEffectiveFrozenDays(inviter *model.User) int { + if inviter != nil && inviter.AffRebateFrozenDays > 0 { + return inviter.AffRebateFrozenDays + } + return GetGlobalRebateFrozenDays() +} + +func ProcessRebateAfterRecharge(invitee *model.User, orderId int, orderType string, orderAmount int) error { + if invitee == nil || invitee.InviterId <= 0 || orderAmount <= 0 { + return nil + } + if !IsRebateEnabled() { + return nil + } + if orderType == "" { + return errors.New("order type is empty") + } + if orderId <= 0 { + return errors.New("order id must be positive") + } + + return model.DB.Transaction(func(tx *gorm.DB) error { + var inviter model.User + if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}). + First(&inviter, "id = ?", invitee.InviterId).Error; err != nil { + return err + } + + var existing model.RebateRecord + err := tx.Where( + "order_type = ? AND order_id = ? AND inviter_id = ? AND invitee_id = ?", + orderType, orderId, inviter.Id, invitee.Id, + ).First(&existing).Error + if err == nil { + return nil + } + if !errors.Is(err, gorm.ErrRecordNotFound) { + return err + } + + rate := GetEffectiveRebateRate(&inviter) + if rate <= 0 { + return nil + } + rebateAmount := orderAmount * rate / 100 + if rebateAmount <= 0 { + return nil + } + + now := time.Now() + record := model.RebateRecord{ + InviterId: inviter.Id, + InviteeId: invitee.Id, + OrderId: orderId, + OrderType: orderType, + OrderAmount: orderAmount, + RebateAmount: rebateAmount, + RatePercent: rate, + Status: "released", + ReleasedAt: &now, + } + if frozenDays := GetEffectiveFrozenDays(&inviter); frozenDays > 0 { + frozenUntil := now.Add(time.Duration(frozenDays) * 24 * time.Hour) + record.Status = "frozen" + record.FrozenUntil = &frozenUntil + record.ReleasedAt = nil + } + + if err := tx.Create(&record).Error; err != nil { + return err + } + + return tx.Model(&model.User{}).Where("id = ?", inviter.Id).Updates(map[string]interface{}{ + "aff_quota": gorm.Expr("aff_quota + ?", rebateAmount), + "aff_history": gorm.Expr("aff_history + ?", rebateAmount), + }).Error + }) +} + +func ReleaseExpiredRebates() error { + return model.DB.Transaction(func(tx *gorm.DB) error { + var records []model.RebateRecord + if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}). + Where("status = ? AND frozen_until IS NOT NULL AND frozen_until <= ?", "frozen", time.Now()). + Order("id ASC"). + Find(&records).Error; err != nil { + return err + } + + releasedAt := time.Now() + for _, record := range records { + result := tx.Model(&model.RebateRecord{}). + Where("id = ? AND status = ?", record.Id, "frozen"). + Updates(map[string]interface{}{ + "status": "released", + "released_at": releasedAt, + }) + if result.Error != nil { + return result.Error + } + } + return nil + }) +} + +func readRebateOption(key string) string { + common.OptionMapRWMutex.RLock() + defer common.OptionMapRWMutex.RUnlock() + return common.OptionMap[key] +} + +func parseNonNegativeRebateInt(value string) int { + parsed, err := strconv.Atoi(value) + if err != nil || parsed < 0 { + return 0 + } + return parsed +} diff --git a/service/rebate_test.go b/service/rebate_test.go new file mode 100644 index 00000000..a6b1c0c7 --- /dev/null +++ b/service/rebate_test.go @@ -0,0 +1,237 @@ +package service + +import ( + "fmt" + "strings" + "testing" + "time" + + "github.com/QuantumNous/new-api/common" + "github.com/QuantumNous/new-api/model" + "github.com/glebarez/sqlite" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +func setupRebateTestDB(t *testing.T) *gorm.DB { + t.Helper() + + previousDB := model.DB + previousLogDB := model.LOG_DB + common.OptionMapRWMutex.RLock() + previousOptionMap := common.OptionMap + common.OptionMapRWMutex.RUnlock() + + common.UsingSQLite = true + common.UsingMySQL = false + common.UsingPostgreSQL = false + common.RedisEnabled = false + + dsn := fmt.Sprintf("file:%s?mode=memory&cache=shared", strings.ReplaceAll(t.Name(), "/", "_")) + db, err := gorm.Open(sqlite.Open(dsn), &gorm.Config{}) + require.NoError(t, err) + require.NoError(t, db.AutoMigrate(&model.User{}, &model.RebateRecord{}, &model.Log{})) + + model.DB = db + model.LOG_DB = db + setRebateOptions(false, "0", "0") + + t.Cleanup(func() { + model.DB = previousDB + model.LOG_DB = previousLogDB + common.OptionMapRWMutex.Lock() + common.OptionMap = previousOptionMap + common.OptionMapRWMutex.Unlock() + sqlDB, err := db.DB() + if err == nil { + _ = sqlDB.Close() + } + }) + + return db +} + +func setRebateOptions(enabled bool, ratePercent string, frozenDays string) { + common.OptionMapRWMutex.Lock() + defer common.OptionMapRWMutex.Unlock() + common.OptionMap = map[string]string{ + model.AffRebateEnabledKey: fmt.Sprintf("%t", enabled), + model.AffRebateRatePercentKey: ratePercent, + model.AffRebateFrozenDaysKey: frozenDays, + } +} + +func createRebateUsers(t *testing.T, db *gorm.DB, inviterOverrides ...func(*model.User)) (*model.User, *model.User) { + t.Helper() + + inviter := &model.User{ + Id: 1, + Username: "inviter_" + strings.ReplaceAll(t.Name(), "/", "_"), + Password: "password123", + AffCode: "aff_inviter_" + strings.ReplaceAll(t.Name(), "/", "_"), + Status: common.UserStatusEnabled, + } + for _, override := range inviterOverrides { + override(inviter) + } + invitee := &model.User{ + Id: 2, + Username: "invitee_" + strings.ReplaceAll(t.Name(), "/", "_"), + Password: "password123", + AffCode: "aff_invitee_" + strings.ReplaceAll(t.Name(), "/", "_"), + Status: common.UserStatusEnabled, + InviterId: inviter.Id, + } + require.NoError(t, db.Create(inviter).Error) + require.NoError(t, db.Create(invitee).Error) + return inviter, invitee +} + +func loadRebateUser(t *testing.T, db *gorm.DB, id int) model.User { + t.Helper() + var user model.User + require.NoError(t, db.First(&user, id).Error) + return user +} + +func TestProcessRebateAfterRechargeSkipsWhenDisabled(t *testing.T) { + db := setupRebateTestDB(t) + _, invitee := createRebateUsers(t, db) + setRebateOptions(false, "10", "7") + + require.NoError(t, ProcessRebateAfterRecharge(invitee, 1001, "topup", 1000)) + + var count int64 + require.NoError(t, db.Model(&model.RebateRecord{}).Count(&count).Error) + require.Equal(t, int64(0), count) + inviter := loadRebateUser(t, db, 1) + require.Equal(t, 0, inviter.AffQuota) + require.Equal(t, 0, inviter.AffHistoryQuota) +} + +func TestProcessRebateAfterRechargeRejectsMissingOrderID(t *testing.T) { + db := setupRebateTestDB(t) + _, invitee := createRebateUsers(t, db) + setRebateOptions(true, "10", "7") + + err := ProcessRebateAfterRecharge(invitee, 0, "topup", 1000) + + require.Error(t, err) + require.Contains(t, err.Error(), "order id") + var count int64 + require.NoError(t, db.Model(&model.RebateRecord{}).Count(&count).Error) + require.Equal(t, int64(0), count) +} + +func TestProcessRebateAfterRechargeUsesGlobalRate(t *testing.T) { + db := setupRebateTestDB(t) + _, invitee := createRebateUsers(t, db) + setRebateOptions(true, "10", "7") + + require.NoError(t, ProcessRebateAfterRecharge(invitee, 1002, "topup", 1000)) + + var record model.RebateRecord + require.NoError(t, db.First(&record).Error) + require.Equal(t, 1, record.InviterId) + require.Equal(t, 2, record.InviteeId) + require.Equal(t, 1002, record.OrderId) + require.Equal(t, "topup", record.OrderType) + require.Equal(t, 1000, record.OrderAmount) + require.Equal(t, 100, record.RebateAmount) + require.Equal(t, 10, record.RatePercent) + require.Equal(t, "frozen", record.Status) + require.NotNil(t, record.FrozenUntil) + require.Nil(t, record.ReleasedAt) + + inviter := loadRebateUser(t, db, 1) + require.Equal(t, 100, inviter.AffQuota) + require.Equal(t, 100, inviter.AffHistoryQuota) +} + +func TestProcessRebateAfterRechargeUsesPerUserOverrides(t *testing.T) { + db := setupRebateTestDB(t) + _, invitee := createRebateUsers(t, db, func(inviter *model.User) { + inviter.AffRebateRatePercent = 25 + inviter.AffRebateFrozenDays = 1 + }) + setRebateOptions(true, "10", "7") + + require.NoError(t, ProcessRebateAfterRecharge(invitee, 1003, "topup", 1000)) + + var record model.RebateRecord + require.NoError(t, db.First(&record).Error) + require.Equal(t, 250, record.RebateAmount) + require.Equal(t, 25, record.RatePercent) + require.Equal(t, "frozen", record.Status) + require.NotNil(t, record.FrozenUntil) + require.WithinDuration(t, time.Now().Add(24*time.Hour), *record.FrozenUntil, time.Minute) + + inviter := loadRebateUser(t, db, 1) + require.Equal(t, 250, inviter.AffQuota) + require.Equal(t, 250, inviter.AffHistoryQuota) +} + +func TestProcessRebateAfterRechargeIsIdempotentForSameOrderAndUsers(t *testing.T) { + db := setupRebateTestDB(t) + _, invitee := createRebateUsers(t, db) + setRebateOptions(true, "10", "0") + + require.NoError(t, ProcessRebateAfterRecharge(invitee, 1004, "topup", 1000)) + require.NoError(t, ProcessRebateAfterRecharge(invitee, 1004, "topup", 1000)) + + var count int64 + require.NoError(t, db.Model(&model.RebateRecord{}).Count(&count).Error) + require.Equal(t, int64(1), count) + inviter := loadRebateUser(t, db, 1) + require.Equal(t, 100, inviter.AffQuota) + require.Equal(t, 100, inviter.AffHistoryQuota) +} + +func TestReleaseExpiredRebatesMarksOnlyDueFrozenRecords(t *testing.T) { + db := setupRebateTestDB(t) + inviter, invitee := createRebateUsers(t, db) + past := time.Now().Add(-time.Hour) + future := time.Now().Add(time.Hour) + require.NoError(t, db.Model(inviter).Updates(map[string]interface{}{ + "aff_quota": 50, + "aff_history": 50, + }).Error) + require.NoError(t, db.Create(&model.RebateRecord{ + InviterId: inviter.Id, + InviteeId: invitee.Id, + OrderId: 1005, + OrderType: "topup", + OrderAmount: 500, + RebateAmount: 50, + RatePercent: 10, + Status: "frozen", + FrozenUntil: &past, + }).Error) + require.NoError(t, db.Create(&model.RebateRecord{ + InviterId: inviter.Id, + InviteeId: invitee.Id, + OrderId: 1006, + OrderType: "topup", + OrderAmount: 500, + RebateAmount: 50, + RatePercent: 10, + Status: "frozen", + FrozenUntil: &future, + }).Error) + + require.NoError(t, ReleaseExpiredRebates()) + + var due model.RebateRecord + require.NoError(t, db.Where("order_id = ?", 1005).First(&due).Error) + require.Equal(t, "released", due.Status) + require.NotNil(t, due.ReleasedAt) + + var notDue model.RebateRecord + require.NoError(t, db.Where("order_id = ?", 1006).First(¬Due).Error) + require.Equal(t, "frozen", notDue.Status) + require.Nil(t, notDue.ReleasedAt) + + reloadedInviter := loadRebateUser(t, db, inviter.Id) + require.Equal(t, 50, reloadedInviter.AffQuota) + require.Equal(t, 50, reloadedInviter.AffHistoryQuota) +}