mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 13:18:18 +08:00
Merge pull request #5224 from wucm667/fix/issue-5190-lock-subscription-renewal
fix(subscription): serialize concurrent renewals
This commit is contained in:
@@ -0,0 +1,43 @@
|
||||
package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/DATA-DOG/go-sqlmock"
|
||||
dbent "github.com/Wei-Shaw/sub2api/ent"
|
||||
_ "github.com/Wei-Shaw/sub2api/ent/runtime"
|
||||
"github.com/Wei-Shaw/sub2api/ent/usersubscription"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"entgo.io/ent/dialect"
|
||||
entsql "entgo.io/ent/dialect/sql"
|
||||
)
|
||||
|
||||
func TestUserSubscriptionGetByIDForUpdateLocksRow(t *testing.T) {
|
||||
var capturedSQL string
|
||||
db, mock, err := sqlmock.New(sqlmock.QueryMatcherOption(captureEntQueryMatcher{actual: &capturedSQL}))
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { _ = db.Close() })
|
||||
|
||||
driver := entsql.OpenDB(dialect.Postgres, db)
|
||||
client := dbent.NewClient(dbent.Driver(driver))
|
||||
t.Cleanup(func() { _ = client.Close() })
|
||||
repo := NewUserSubscriptionRepository(client)
|
||||
now := time.Date(2026, 8, 2, 12, 0, 0, 0, time.UTC)
|
||||
|
||||
mock.ExpectQuery("locked subscription").WillReturnRows(
|
||||
sqlmock.NewRows(usersubscription.Columns).AddRow(
|
||||
int64(7), now, now, nil, int64(11), int64(13), now, now.AddDate(0, 0, 30), "active",
|
||||
nil, nil, nil, 0.0, 0.0, 0.0, nil, now, "renewal",
|
||||
),
|
||||
)
|
||||
|
||||
sub, err := repo.GetByIDForUpdate(context.Background(), 7)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, int64(7), sub.ID)
|
||||
require.NoError(t, mock.ExpectationsWereMet())
|
||||
require.Contains(t, strings.ToUpper(normalizeSQLWhitespace(capturedSQL)), "FOR UPDATE")
|
||||
}
|
||||
@@ -75,6 +75,18 @@ func (r *userSubscriptionRepository) GetByID(ctx context.Context, id int64) (*se
|
||||
return userSubscriptionEntityToService(m), nil
|
||||
}
|
||||
|
||||
func (r *userSubscriptionRepository) GetByIDForUpdate(ctx context.Context, id int64) (*service.UserSubscription, error) {
|
||||
client := clientFromContext(ctx, r.client)
|
||||
m, err := client.UserSubscription.Query().
|
||||
Where(usersubscription.IDEQ(id)).
|
||||
ForUpdate().
|
||||
Only(ctx)
|
||||
if err != nil {
|
||||
return nil, translatePersistenceError(err, service.ErrSubscriptionNotFound, nil)
|
||||
}
|
||||
return userSubscriptionEntityToService(m), nil
|
||||
}
|
||||
|
||||
func (r *userSubscriptionRepository) GetByIDIncludeDeleted(ctx context.Context, id int64) (*service.UserSubscription, error) {
|
||||
client := clientFromContext(ctx, r.client)
|
||||
queryCtx := mixins.SkipSoftDelete(ctx)
|
||||
|
||||
@@ -2169,6 +2169,9 @@ func (stubUserSubscriptionRepo) Create(ctx context.Context, sub *service.UserSub
|
||||
func (stubUserSubscriptionRepo) GetByID(ctx context.Context, id int64) (*service.UserSubscription, error) {
|
||||
return nil, errors.New("not implemented")
|
||||
}
|
||||
func (stubUserSubscriptionRepo) GetByIDForUpdate(ctx context.Context, id int64) (*service.UserSubscription, error) {
|
||||
return nil, errors.New("not implemented")
|
||||
}
|
||||
func (stubUserSubscriptionRepo) GetByIDIncludeDeleted(ctx context.Context, id int64) (*service.UserSubscription, error) {
|
||||
return nil, errors.New("not implemented")
|
||||
}
|
||||
|
||||
@@ -176,6 +176,10 @@ func (f fakeGoogleSubscriptionRepo) GetByID(ctx context.Context, id int64) (*ser
|
||||
}
|
||||
return nil, errors.New("not implemented")
|
||||
}
|
||||
|
||||
func (f fakeGoogleSubscriptionRepo) GetByIDForUpdate(ctx context.Context, id int64) (*service.UserSubscription, error) {
|
||||
return f.GetByID(ctx, id)
|
||||
}
|
||||
func (f fakeGoogleSubscriptionRepo) GetByIDIncludeDeleted(ctx context.Context, id int64) (*service.UserSubscription, error) {
|
||||
return nil, errors.New("not implemented")
|
||||
}
|
||||
|
||||
@@ -1683,6 +1683,10 @@ func (r *stubUserSubscriptionRepo) GetByID(ctx context.Context, id int64) (*serv
|
||||
return nil, errors.New("not implemented")
|
||||
}
|
||||
|
||||
func (r *stubUserSubscriptionRepo) GetByIDForUpdate(ctx context.Context, id int64) (*service.UserSubscription, error) {
|
||||
return r.GetByID(ctx, id)
|
||||
}
|
||||
|
||||
func (r *stubUserSubscriptionRepo) GetByIDIncludeDeleted(ctx context.Context, id int64) (*service.UserSubscription, error) {
|
||||
return nil, errors.New("not implemented")
|
||||
}
|
||||
|
||||
@@ -111,6 +111,9 @@ func (userSubRepoNoop) Create(context.Context, *UserSubscription) error {
|
||||
func (userSubRepoNoop) GetByID(context.Context, int64) (*UserSubscription, error) {
|
||||
panic("unexpected GetByID call")
|
||||
}
|
||||
func (userSubRepoNoop) GetByIDForUpdate(context.Context, int64) (*UserSubscription, error) {
|
||||
panic("unexpected GetByIDForUpdate call")
|
||||
}
|
||||
func (userSubRepoNoop) GetByIDIncludeDeleted(context.Context, int64) (*UserSubscription, error) {
|
||||
panic("unexpected GetByIDIncludeDeleted call")
|
||||
}
|
||||
@@ -249,6 +252,10 @@ func (s *subscriptionUserSubRepoStub) GetByID(_ context.Context, id int64) (*Use
|
||||
return &cp, nil
|
||||
}
|
||||
|
||||
func (s *subscriptionUserSubRepoStub) GetByIDForUpdate(ctx context.Context, id int64) (*UserSubscription, error) {
|
||||
return s.GetByID(ctx, id)
|
||||
}
|
||||
|
||||
func (s *subscriptionUserSubRepoStub) Update(_ context.Context, sub *UserSubscription) error {
|
||||
if sub == nil {
|
||||
return ErrSubscriptionNilInput
|
||||
|
||||
@@ -22,6 +22,10 @@ func (r *subscriptionExpiryRepoStub) GetByID(context.Context, int64) (*UserSubsc
|
||||
return nil, ErrSubscriptionNotFound
|
||||
}
|
||||
|
||||
func (r *subscriptionExpiryRepoStub) GetByIDForUpdate(context.Context, int64) (*UserSubscription, error) {
|
||||
return nil, ErrSubscriptionNotFound
|
||||
}
|
||||
|
||||
func (r *subscriptionExpiryRepoStub) GetByIDIncludeDeleted(context.Context, int64) (*UserSubscription, error) {
|
||||
return nil, ErrSubscriptionNotFound
|
||||
}
|
||||
|
||||
@@ -0,0 +1,143 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
type lockingRenewalRepo struct {
|
||||
userSubRepoNoop
|
||||
mu sync.Mutex
|
||||
stale UserSubscription
|
||||
current UserSubscription
|
||||
lockReads int
|
||||
}
|
||||
|
||||
func (r *lockingRenewalRepo) ExistsByUserIDAndGroupID(context.Context, int64, int64) (bool, error) {
|
||||
return true, nil
|
||||
}
|
||||
|
||||
func (r *lockingRenewalRepo) GetByUserIDAndGroupID(context.Context, int64, int64) (*UserSubscription, error) {
|
||||
copy := r.stale
|
||||
return ©, nil
|
||||
}
|
||||
|
||||
func (r *lockingRenewalRepo) GetByID(_ context.Context, _ int64) (*UserSubscription, error) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
copy := r.current
|
||||
return ©, nil
|
||||
}
|
||||
|
||||
func (r *lockingRenewalRepo) GetByIDForUpdate(_ context.Context, _ int64) (*UserSubscription, error) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
r.lockReads++
|
||||
copy := r.current
|
||||
return ©, nil
|
||||
}
|
||||
|
||||
func (r *lockingRenewalRepo) ExtendExpiry(_ context.Context, _ int64, expiresAt time.Time) error {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
r.current.ExpiresAt = expiresAt
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *lockingRenewalRepo) UpdateStatus(_ context.Context, _ int64, status string) error {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
r.current.Status = status
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *lockingRenewalRepo) UpdateNotes(_ context.Context, _ int64, notes string) error {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
r.current.Notes = notes
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *lockingRenewalRepo) Update(_ context.Context, sub *UserSubscription) error {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
r.current = *sub
|
||||
return nil
|
||||
}
|
||||
|
||||
func TestAssignOrExtendSubscriptionUsesLockedCurrentRow(t *testing.T) {
|
||||
now := time.Date(2026, 8, 2, 12, 0, 0, 0, time.UTC)
|
||||
lockedExpiry := now.AddDate(0, 0, 20)
|
||||
windowStart := now.Add(-24 * time.Hour)
|
||||
repo := &lockingRenewalRepo{
|
||||
stale: UserSubscription{ID: 7, UserID: 11, GroupID: 13, ExpiresAt: now.Add(-time.Hour), Status: SubscriptionStatusExpired, Notes: "stale"},
|
||||
current: UserSubscription{
|
||||
ID: 7, UserID: 11, GroupID: 13, StartsAt: now.AddDate(0, 0, -10), ExpiresAt: lockedExpiry,
|
||||
Status: SubscriptionStatusSuspended, Notes: "current", DailyWindowStart: &windowStart, DailyUsageUSD: 4,
|
||||
},
|
||||
}
|
||||
svc := NewSubscriptionService(&subscriptionGroupRepoStub{group: &Group{ID: 13, SubscriptionType: SubscriptionTypeSubscription}}, repo, nil, nil, nil)
|
||||
svc.now = func() time.Time { return now }
|
||||
|
||||
sub, extended, err := svc.AssignOrExtendSubscription(context.Background(), &AssignSubscriptionInput{
|
||||
UserID: 11, GroupID: 13, ValidityDays: 5, Notes: "renewed",
|
||||
})
|
||||
|
||||
require.NoError(t, err)
|
||||
require.True(t, extended)
|
||||
require.Equal(t, 1, repo.lockReads)
|
||||
require.Equal(t, lockedExpiry.AddDate(0, 0, 5), sub.ExpiresAt)
|
||||
require.Equal(t, SubscriptionStatusActive, sub.Status)
|
||||
require.Equal(t, "current\nrenewed", sub.Notes)
|
||||
require.Equal(t, windowStart, *sub.DailyWindowStart)
|
||||
require.Equal(t, float64(4), sub.DailyUsageUSD)
|
||||
}
|
||||
|
||||
func TestAssignOrExtendSubscriptionSerializedRenewalsAccumulateDays(t *testing.T) {
|
||||
now := time.Date(2026, 8, 2, 12, 0, 0, 0, time.UTC)
|
||||
initialExpiry := now.AddDate(0, 0, 10)
|
||||
stale := UserSubscription{ID: 17, UserID: 21, GroupID: 23, StartsAt: now, ExpiresAt: initialExpiry, Status: SubscriptionStatusActive}
|
||||
repo := &lockingRenewalRepo{stale: stale, current: stale}
|
||||
svc := NewSubscriptionService(&subscriptionGroupRepoStub{group: &Group{ID: 23, SubscriptionType: SubscriptionTypeSubscription}}, repo, nil, nil, nil)
|
||||
svc.now = func() time.Time { return now }
|
||||
input := &AssignSubscriptionInput{UserID: 21, GroupID: 23, ValidityDays: 7}
|
||||
|
||||
_, _, err := svc.AssignOrExtendSubscription(context.Background(), input)
|
||||
require.NoError(t, err)
|
||||
second, _, err := svc.AssignOrExtendSubscription(context.Background(), input)
|
||||
require.NoError(t, err)
|
||||
|
||||
require.Equal(t, 2, repo.lockReads)
|
||||
require.Equal(t, initialExpiry.AddDate(0, 0, 14), second.ExpiresAt)
|
||||
}
|
||||
|
||||
func TestAssignSubscriptionDoesNotReactivateRowSuspendedAfterStaleRead(t *testing.T) {
|
||||
now := time.Date(2026, 8, 2, 12, 0, 0, 0, time.UTC)
|
||||
windowStart := now.Add(-24 * time.Hour)
|
||||
current := UserSubscription{
|
||||
ID: 27, UserID: 31, GroupID: 33, StartsAt: now.AddDate(0, 0, -10), ExpiresAt: now.Add(-time.Hour),
|
||||
Status: SubscriptionStatusSuspended, Notes: "suspended", DailyWindowStart: &windowStart, DailyUsageUSD: 4,
|
||||
}
|
||||
repo := &lockingRenewalRepo{
|
||||
stale: UserSubscription{ID: 27, UserID: 31, GroupID: 33, ExpiresAt: now.Add(-time.Hour), Status: SubscriptionStatusExpired},
|
||||
current: current,
|
||||
}
|
||||
svc := NewSubscriptionService(&subscriptionGroupRepoStub{group: &Group{ID: 33, SubscriptionType: SubscriptionTypeSubscription}}, repo, nil, nil, nil)
|
||||
svc.now = func() time.Time { return now }
|
||||
|
||||
sub, reused, err := svc.assignSubscriptionWithReuse(context.Background(), &AssignSubscriptionInput{
|
||||
UserID: 31, GroupID: 33, ValidityDays: 5, Notes: "renewed",
|
||||
})
|
||||
|
||||
require.NoError(t, err)
|
||||
require.True(t, reused)
|
||||
require.Equal(t, 1, repo.lockReads)
|
||||
require.Equal(t, current, repo.current)
|
||||
require.Equal(t, SubscriptionStatusSuspended, sub.Status)
|
||||
require.Equal(t, current.ExpiresAt, sub.ExpiresAt)
|
||||
require.Equal(t, current.Notes, sub.Notes)
|
||||
}
|
||||
@@ -244,24 +244,7 @@ func (s *SubscriptionService) assignOrExtendSubscription(ctx context.Context, in
|
||||
|
||||
// 已有订阅,执行续期(在事务中完成所有更新)
|
||||
if existingSub != nil {
|
||||
now := time.Now()
|
||||
var newExpiresAt time.Time
|
||||
|
||||
isExpired := !existingSub.ExpiresAt.After(now)
|
||||
if !isExpired {
|
||||
// 未过期:从当前过期时间累加
|
||||
newExpiresAt = existingSub.ExpiresAt.AddDate(0, 0, validityDays)
|
||||
} else {
|
||||
// 已过期:从当前时间开始计算
|
||||
newExpiresAt = now.AddDate(0, 0, validityDays)
|
||||
}
|
||||
|
||||
// 确保不超过最大过期时间
|
||||
if newExpiresAt.After(MaxExpiresAt) {
|
||||
newExpiresAt = MaxExpiresAt
|
||||
}
|
||||
|
||||
if err := s.updateExistingSubscriptionTerm(ctx, existingSub, input.Notes, now, newExpiresAt, isExpired); err != nil {
|
||||
if err := s.updateExistingSubscriptionTerm(ctx, existingSub.ID, validityDays, input.Notes, false); err != nil {
|
||||
return nil, false, err
|
||||
}
|
||||
|
||||
@@ -305,15 +288,42 @@ func (s *SubscriptionService) maybeInvalidateAssignmentCaches(userID, groupID in
|
||||
|
||||
func (s *SubscriptionService) updateExistingSubscriptionTerm(
|
||||
ctx context.Context,
|
||||
existingSub *UserSubscription,
|
||||
subscriptionID int64,
|
||||
validityDays int,
|
||||
notes string,
|
||||
startsAt time.Time,
|
||||
newExpiresAt time.Time,
|
||||
isExpired bool,
|
||||
assignmentSemantics bool,
|
||||
) error {
|
||||
return s.withSubscriptionUpdateTx(ctx, func(txCtx context.Context) error {
|
||||
existingSub, err := s.userSubRepo.GetByIDForUpdate(txCtx, subscriptionID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("lock subscription for renewal: %w", err)
|
||||
}
|
||||
if assignmentSemantics && existingSub.Status == SubscriptionStatusSuspended {
|
||||
return nil
|
||||
}
|
||||
|
||||
now := time.Now()
|
||||
if s.now != nil {
|
||||
now = s.now()
|
||||
}
|
||||
isExpired := !existingSub.ExpiresAt.After(now)
|
||||
if assignmentSemantics {
|
||||
isExpired = existingSub.Status == SubscriptionStatusExpired ||
|
||||
(existingSub.Status != SubscriptionStatusSuspended && !existingSub.ExpiresAt.After(now))
|
||||
}
|
||||
newExpiresAt := existingSub.ExpiresAt.AddDate(0, 0, validityDays)
|
||||
if isExpired {
|
||||
renewed := renewedSubscriptionTerm(existingSub, notes, startsAt, newExpiresAt)
|
||||
newExpiresAt = now.AddDate(0, 0, validityDays)
|
||||
}
|
||||
if newExpiresAt.After(MaxExpiresAt) {
|
||||
newExpiresAt = MaxExpiresAt
|
||||
}
|
||||
if assignmentSemantics && strings.TrimSpace(existingSub.Notes) == strings.TrimSpace(notes) {
|
||||
notes = ""
|
||||
}
|
||||
|
||||
if isExpired {
|
||||
renewed := renewedSubscriptionTerm(existingSub, notes, now, newExpiresAt)
|
||||
if err := s.userSubRepo.Update(txCtx, renewed); err != nil {
|
||||
return fmt.Errorf("renew expired subscription: %w", err)
|
||||
}
|
||||
@@ -514,15 +524,7 @@ func (s *SubscriptionService) assignSubscriptionWithReuse(ctx context.Context, i
|
||||
if sub.Status == SubscriptionStatusExpired ||
|
||||
(sub.Status != SubscriptionStatusSuspended && !sub.ExpiresAt.After(now)) {
|
||||
validityDays := normalizeAssignValidityDays(input.ValidityDays)
|
||||
newExpiresAt := now.AddDate(0, 0, validityDays)
|
||||
if newExpiresAt.After(MaxExpiresAt) {
|
||||
newExpiresAt = MaxExpiresAt
|
||||
}
|
||||
renewalNotes := input.Notes
|
||||
if strings.TrimSpace(sub.Notes) == strings.TrimSpace(input.Notes) {
|
||||
renewalNotes = ""
|
||||
}
|
||||
if err := s.updateExistingSubscriptionTerm(ctx, sub, renewalNotes, now, newExpiresAt, true); err != nil {
|
||||
if err := s.updateExistingSubscriptionTerm(ctx, sub.ID, validityDays, input.Notes, true); err != nil {
|
||||
return nil, false, err
|
||||
}
|
||||
s.maybeInvalidateAssignmentCaches(input.UserID, input.GroupID, false)
|
||||
|
||||
@@ -10,6 +10,7 @@ import (
|
||||
type UserSubscriptionRepository interface {
|
||||
Create(ctx context.Context, sub *UserSubscription) error
|
||||
GetByID(ctx context.Context, id int64) (*UserSubscription, error)
|
||||
GetByIDForUpdate(ctx context.Context, id int64) (*UserSubscription, error)
|
||||
GetByIDIncludeDeleted(ctx context.Context, id int64) (*UserSubscription, error)
|
||||
GetByUserIDAndGroupID(ctx context.Context, userID, groupID int64) (*UserSubscription, error)
|
||||
GetActiveByUserIDAndGroupID(ctx context.Context, userID, groupID int64) (*UserSubscription, error)
|
||||
|
||||
Reference in New Issue
Block a user