This commit is contained in:
@@ -131,9 +131,8 @@ func (l *CloseOrderLogic) CloseOrder(req *types.CloseOrderRequest) error {
|
||||
)
|
||||
return err
|
||||
}
|
||||
// update user cache
|
||||
return l.svcCtx.UserModel.UpdateUserCache(l.ctx, userInfo)
|
||||
}
|
||||
// Note: user cache will be updated after transaction commits
|
||||
if sub.Inventory != -1 {
|
||||
sub.Inventory++
|
||||
if e := l.svcCtx.SubscribeModel.Update(l.ctx, sub, tx); e != nil {
|
||||
@@ -151,6 +150,19 @@ func (l *CloseOrderLogic) CloseOrder(req *types.CloseOrderRequest) error {
|
||||
logger.Errorf("[CloseOrder] Transaction failed: %v", err.Error())
|
||||
return err
|
||||
}
|
||||
|
||||
// Update user cache after transaction commits successfully
|
||||
if orderInfo.GiftAmount > 0 && orderInfo.UserId != 0 {
|
||||
if userInfo, findErr := l.svcCtx.UserModel.FindOne(l.ctx, orderInfo.UserId); findErr == nil {
|
||||
if clearErr := l.svcCtx.UserModel.ClearUserCache(l.ctx, userInfo); clearErr != nil {
|
||||
l.Errorw("[CloseOrder] failed to clear user cache",
|
||||
logger.Field("error", clearErr.Error()),
|
||||
logger.Field("user_id", orderInfo.UserId),
|
||||
)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"math"
|
||||
"time"
|
||||
|
||||
commonLogic "github.com/perfect-panel/server/internal/logic/common"
|
||||
"github.com/perfect-panel/server/internal/model/order"
|
||||
@@ -108,6 +109,23 @@ func (l *PreCreateOrderLogic) PreCreateOrder(req *types.PurchaseOrderRequest) (r
|
||||
}
|
||||
}
|
||||
|
||||
// check new user only restriction
|
||||
if !isSingleModeRenewal && sub.NewUserOnly != nil && *sub.NewUserOnly {
|
||||
if time.Since(u.CreatedAt) > 24*time.Hour {
|
||||
return nil, errors.Wrapf(xerr.NewErrCode(xerr.SubscribeNewUserOnly), "not a new user")
|
||||
}
|
||||
var historyCount int64
|
||||
if e := l.svcCtx.DB.Model(&order.Order{}).
|
||||
Where("user_id = ? AND subscribe_id = ? AND type = 1 AND status IN ?",
|
||||
u.Id, targetSubscribeID, []uint8{2, 5}).
|
||||
Count(&historyCount).Error; e != nil {
|
||||
return nil, errors.Wrapf(xerr.NewErrCode(xerr.DatabaseQueryError), "check new user purchase history error: %v", e.Error())
|
||||
}
|
||||
if historyCount >= 1 {
|
||||
return nil, errors.Wrapf(xerr.NewErrCode(xerr.SubscribeNewUserOnly), "already purchased new user plan")
|
||||
}
|
||||
}
|
||||
|
||||
var discount float64 = 1
|
||||
if sub.Discount != "" {
|
||||
var dis []types.SubscribeDiscount
|
||||
|
||||
@@ -270,6 +270,23 @@ func (l *PurchaseLogic) Purchase(req *types.PurchaseOrderRequest) (resp *types.P
|
||||
}
|
||||
}
|
||||
|
||||
// check new user only restriction inside transaction to prevent race condition
|
||||
if orderInfo.Type == 1 && sub.NewUserOnly != nil && *sub.NewUserOnly {
|
||||
if time.Since(u.CreatedAt) > 24*time.Hour {
|
||||
return errors.Wrapf(xerr.NewErrCode(xerr.SubscribeNewUserOnly), "not a new user")
|
||||
}
|
||||
var historyCount int64
|
||||
if e := db.Model(&order.Order{}).
|
||||
Where("user_id = ? AND subscribe_id = ? AND type = 1 AND status IN ?",
|
||||
u.Id, targetSubscribeID, []int{2, 5}).
|
||||
Count(&historyCount).Error; e != nil {
|
||||
return errors.Wrapf(xerr.NewErrCode(xerr.DatabaseQueryError), "check new user purchase history error: %v", e.Error())
|
||||
}
|
||||
if historyCount >= 1 {
|
||||
return errors.Wrapf(xerr.NewErrCode(xerr.SubscribeNewUserOnly), "already purchased new user plan")
|
||||
}
|
||||
}
|
||||
|
||||
// update user gift amount and create deduction record
|
||||
if orderInfo.GiftAmount > 0 {
|
||||
// deduct gift amount from user
|
||||
@@ -319,7 +336,11 @@ func (l *PurchaseLogic) Purchase(req *types.PurchaseOrderRequest) (resp *types.P
|
||||
})
|
||||
if err != nil {
|
||||
l.Errorw("[Purchase] Database insert error", logger.Field("error", err.Error()), logger.Field("orderInfo", orderInfo))
|
||||
|
||||
// Propagate business errors (e.g. SubscribeNewUserOnly, SubscribeQuotaLimit) directly.
|
||||
var codeErr *xerr.CodeError
|
||||
if errors.As(err, &codeErr) {
|
||||
return nil, err
|
||||
}
|
||||
return nil, errors.Wrapf(xerr.NewErrCode(xerr.DatabaseInsertError), "insert order error: %v", err.Error())
|
||||
}
|
||||
// Deferred task
|
||||
|
||||
@@ -1,33 +0,0 @@
|
||||
package subscribe
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
commonLogic "github.com/perfect-panel/server/internal/logic/common"
|
||||
"github.com/perfect-panel/server/internal/types"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestFillUserSubscribeInfoEntitlementFields(t *testing.T) {
|
||||
sub := &types.UserSubscribeInfo{}
|
||||
entitlement := &commonLogic.EntitlementContext{
|
||||
EffectiveUserID: 3001,
|
||||
Source: commonLogic.EntitlementSourceFamilyOwner,
|
||||
OwnerUserID: 3001,
|
||||
ReadOnly: true,
|
||||
}
|
||||
|
||||
fillUserSubscribeInfoEntitlementFields(sub, entitlement)
|
||||
|
||||
require.Equal(t, commonLogic.EntitlementSourceFamilyOwner, sub.EntitlementSource)
|
||||
require.Equal(t, int64(3001), sub.EntitlementOwnerUserId)
|
||||
require.True(t, sub.ReadOnly)
|
||||
}
|
||||
|
||||
func TestNormalizeSubscribeNodeTags(t *testing.T) {
|
||||
tags := normalizeSubscribeNodeTags("美国, 日本, , 美国, ,日本")
|
||||
require.Equal(t, []string{"美国", "日本"}, tags)
|
||||
|
||||
empty := normalizeSubscribeNodeTags("")
|
||||
require.Nil(t, empty)
|
||||
}
|
||||
@@ -45,6 +45,9 @@ func (h *accountMergeHelper) mergeIntoOwner(ownerUserID, deviceUserID int64, sou
|
||||
DeviceUserID: deviceUserID,
|
||||
}
|
||||
|
||||
// Capture device user's auth methods BEFORE the transaction migrates them
|
||||
deviceAuthMethods, _ := h.svcCtx.UserModel.FindUserAuthMethods(h.ctx, deviceUserID)
|
||||
|
||||
err := h.svcCtx.DB.WithContext(h.ctx).Transaction(func(tx *gorm.DB) error {
|
||||
var owner modelUser.User
|
||||
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).
|
||||
@@ -114,7 +117,7 @@ func (h *accountMergeHelper) mergeIntoOwner(ownerUserID, deviceUserID int64, sou
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if err := h.clearCaches(result); err != nil {
|
||||
if err := h.clearCaches(result, deviceAuthMethods); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -129,16 +132,32 @@ func (h *accountMergeHelper) mergeIntoOwner(ownerUserID, deviceUserID int64, sou
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (h *accountMergeHelper) clearCaches(result *accountMergeResult) error {
|
||||
func (h *accountMergeHelper) clearCaches(result *accountMergeResult, deviceAuthMethods []*modelUser.AuthMethods) error {
|
||||
if result == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
if err := h.svcCtx.UserModel.ClearUserCache(h.ctx,
|
||||
&modelUser.User{Id: result.OwnerUserID},
|
||||
&modelUser.User{Id: result.DeviceUserID},
|
||||
); err != nil {
|
||||
return err
|
||||
// Fetch owner user with AuthMethods for proper cache key generation
|
||||
var users []*modelUser.User
|
||||
if u, err := h.svcCtx.UserModel.FindOne(h.ctx, result.OwnerUserID); err == nil {
|
||||
users = append(users, u)
|
||||
}
|
||||
// For device user, FindOne won't have AuthMethods anymore (migrated in tx),
|
||||
// so we build a minimal User with the pre-captured auth methods
|
||||
deviceUser := &modelUser.User{Id: result.DeviceUserID}
|
||||
if len(deviceAuthMethods) > 0 {
|
||||
authMethods := make([]modelUser.AuthMethods, len(deviceAuthMethods))
|
||||
for i, am := range deviceAuthMethods {
|
||||
authMethods[i] = *am
|
||||
}
|
||||
deviceUser.AuthMethods = authMethods
|
||||
}
|
||||
users = append(users, deviceUser)
|
||||
|
||||
if len(users) > 0 {
|
||||
if err := h.svcCtx.UserModel.ClearUserCache(h.ctx, users...); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
if len(result.MovedDevices) > 0 {
|
||||
|
||||
@@ -1,109 +0,0 @@
|
||||
package user
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/alicebob/miniredis/v2"
|
||||
"github.com/perfect-panel/server/internal/svc"
|
||||
"github.com/redis/go-redis/v9"
|
||||
)
|
||||
|
||||
func TestClearAllSessions_RemovesUserSessionsAndDeviceMappings(t *testing.T) {
|
||||
logic, redisClient, cleanup := newDeleteAccountRedisTestLogic(t)
|
||||
defer cleanup()
|
||||
|
||||
mustRedisSet(t, redisClient, "auth:session_id:sid-user-1", "1001")
|
||||
mustRedisSet(t, redisClient, "auth:session_id:sid-user-2", "1001")
|
||||
mustRedisSet(t, redisClient, "auth:session_id:sid-other", "2002")
|
||||
|
||||
mustRedisSet(t, redisClient, "auth:session_id:detail:sid-user-1", "detail")
|
||||
mustRedisSet(t, redisClient, "auth:session_id:detail:sid-other", "detail")
|
||||
|
||||
mustRedisSet(t, redisClient, "auth:device_identifier:dev-user-1", "sid-user-1")
|
||||
mustRedisSet(t, redisClient, "auth:device_identifier:dev-user-2", "sid-user-2")
|
||||
mustRedisSet(t, redisClient, "auth:device_identifier:dev-other", "sid-other")
|
||||
|
||||
mustRedisZAdd(t, redisClient, "auth:user_sessions:1001", "sid-user-3", 1)
|
||||
mustRedisSet(t, redisClient, "auth:session_id:sid-user-3", "1001")
|
||||
|
||||
logic.clearAllSessions(1001)
|
||||
|
||||
mustRedisNotExist(t, redisClient, "auth:session_id:sid-user-1")
|
||||
mustRedisNotExist(t, redisClient, "auth:session_id:sid-user-2")
|
||||
mustRedisNotExist(t, redisClient, "auth:session_id:sid-user-3")
|
||||
mustRedisNotExist(t, redisClient, "auth:session_id:detail:sid-user-1")
|
||||
mustRedisNotExist(t, redisClient, "auth:user_sessions:1001")
|
||||
mustRedisNotExist(t, redisClient, "auth:device_identifier:dev-user-1")
|
||||
mustRedisNotExist(t, redisClient, "auth:device_identifier:dev-user-2")
|
||||
|
||||
mustRedisExist(t, redisClient, "auth:session_id:sid-other")
|
||||
mustRedisExist(t, redisClient, "auth:session_id:detail:sid-other")
|
||||
mustRedisExist(t, redisClient, "auth:device_identifier:dev-other")
|
||||
}
|
||||
|
||||
func TestClearAllSessions_ScanFallbackWorksWithoutUserSessionIndex(t *testing.T) {
|
||||
logic, redisClient, cleanup := newDeleteAccountRedisTestLogic(t)
|
||||
defer cleanup()
|
||||
|
||||
mustRedisSet(t, redisClient, "auth:session_id:sid-a", "3003")
|
||||
mustRedisSet(t, redisClient, "auth:session_id:sid-b", "3003")
|
||||
mustRedisSet(t, redisClient, "auth:session_id:sid-c", "4004")
|
||||
|
||||
logic.clearAllSessions(3003)
|
||||
|
||||
mustRedisNotExist(t, redisClient, "auth:session_id:sid-a")
|
||||
mustRedisNotExist(t, redisClient, "auth:session_id:sid-b")
|
||||
mustRedisExist(t, redisClient, "auth:session_id:sid-c")
|
||||
}
|
||||
|
||||
func newDeleteAccountRedisTestLogic(t *testing.T) (*DeleteAccountLogic, *redis.Client, func()) {
|
||||
t.Helper()
|
||||
|
||||
miniRedis := miniredis.RunT(t)
|
||||
redisClient := redis.NewClient(&redis.Options{Addr: miniRedis.Addr()})
|
||||
logic := NewDeleteAccountLogic(context.Background(), &svc.ServiceContext{Redis: redisClient})
|
||||
|
||||
cleanup := func() {
|
||||
_ = redisClient.Close()
|
||||
miniRedis.Close()
|
||||
}
|
||||
return logic, redisClient, cleanup
|
||||
}
|
||||
|
||||
func mustRedisSet(t *testing.T, redisClient *redis.Client, key, value string) {
|
||||
t.Helper()
|
||||
if err := redisClient.Set(context.Background(), key, value, time.Hour).Err(); err != nil {
|
||||
t.Fatalf("redis set %s failed: %v", key, err)
|
||||
}
|
||||
}
|
||||
|
||||
func mustRedisZAdd(t *testing.T, redisClient *redis.Client, key, member string, score float64) {
|
||||
t.Helper()
|
||||
if err := redisClient.ZAdd(context.Background(), key, redis.Z{Member: member, Score: score}).Err(); err != nil {
|
||||
t.Fatalf("redis zadd %s failed: %v", key, err)
|
||||
}
|
||||
}
|
||||
|
||||
func mustRedisExist(t *testing.T, redisClient *redis.Client, key string) {
|
||||
t.Helper()
|
||||
exists, err := redisClient.Exists(context.Background(), key).Result()
|
||||
if err != nil {
|
||||
t.Fatalf("redis exists %s failed: %v", key, err)
|
||||
}
|
||||
if exists == 0 {
|
||||
t.Fatalf("expected redis key %s to exist", key)
|
||||
}
|
||||
}
|
||||
|
||||
func mustRedisNotExist(t *testing.T, redisClient *redis.Client, key string) {
|
||||
t.Helper()
|
||||
exists, err := redisClient.Exists(context.Background(), key).Result()
|
||||
if err != nil {
|
||||
t.Fatalf("redis exists %s failed: %v", key, err)
|
||||
}
|
||||
if exists != 0 {
|
||||
t.Fatalf("expected redis key %s to be deleted", key)
|
||||
}
|
||||
}
|
||||
@@ -1,128 +0,0 @@
|
||||
package user
|
||||
|
||||
import (
|
||||
stderrors "errors"
|
||||
"testing"
|
||||
|
||||
modelUser "github.com/perfect-panel/server/internal/model/user"
|
||||
"github.com/perfect-panel/server/pkg/xerr"
|
||||
pkgerrors "github.com/pkg/errors"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func extractFamilyJoinCode(err error) uint32 {
|
||||
if err == nil {
|
||||
return 0
|
||||
}
|
||||
|
||||
var codeErr *xerr.CodeError
|
||||
if stderrors.As(pkgerrors.Cause(err), &codeErr) {
|
||||
return codeErr.GetErrCode()
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func TestValidateMemberJoinConflict(t *testing.T) {
|
||||
ownerFamilyID := int64(11)
|
||||
|
||||
testCases := []struct {
|
||||
name string
|
||||
ownerFamily int64
|
||||
memberRecord *modelUser.UserFamilyMember
|
||||
wantCode uint32
|
||||
}{
|
||||
{
|
||||
name: "no member record",
|
||||
ownerFamily: ownerFamilyID,
|
||||
wantCode: 0,
|
||||
},
|
||||
{
|
||||
name: "same family active member",
|
||||
ownerFamily: ownerFamilyID,
|
||||
memberRecord: &modelUser.UserFamilyMember{
|
||||
FamilyId: ownerFamilyID,
|
||||
Status: modelUser.FamilyMemberActive,
|
||||
},
|
||||
wantCode: xerr.FamilyAlreadyBound,
|
||||
},
|
||||
{
|
||||
name: "same family left member",
|
||||
ownerFamily: ownerFamilyID,
|
||||
memberRecord: &modelUser.UserFamilyMember{
|
||||
FamilyId: ownerFamilyID,
|
||||
Status: modelUser.FamilyMemberLeft,
|
||||
},
|
||||
wantCode: 0,
|
||||
},
|
||||
{
|
||||
name: "same family removed member",
|
||||
ownerFamily: ownerFamilyID,
|
||||
memberRecord: &modelUser.UserFamilyMember{
|
||||
FamilyId: ownerFamilyID,
|
||||
Status: modelUser.FamilyMemberRemoved,
|
||||
},
|
||||
wantCode: 0,
|
||||
},
|
||||
{
|
||||
name: "cross family active member",
|
||||
ownerFamily: ownerFamilyID,
|
||||
memberRecord: &modelUser.UserFamilyMember{
|
||||
FamilyId: ownerFamilyID + 1,
|
||||
Status: modelUser.FamilyMemberActive,
|
||||
},
|
||||
wantCode: xerr.FamilyCrossBindForbidden,
|
||||
},
|
||||
{
|
||||
name: "cross family left member",
|
||||
ownerFamily: ownerFamilyID,
|
||||
memberRecord: &modelUser.UserFamilyMember{
|
||||
FamilyId: ownerFamilyID + 1,
|
||||
Status: modelUser.FamilyMemberLeft,
|
||||
},
|
||||
wantCode: 0,
|
||||
},
|
||||
{
|
||||
name: "cross family removed member",
|
||||
ownerFamily: ownerFamilyID,
|
||||
memberRecord: &modelUser.UserFamilyMember{
|
||||
FamilyId: ownerFamilyID + 1,
|
||||
Status: modelUser.FamilyMemberRemoved,
|
||||
},
|
||||
wantCode: 0,
|
||||
},
|
||||
}
|
||||
|
||||
for _, testCase := range testCases {
|
||||
t.Run(testCase.name, func(t *testing.T) {
|
||||
err := validateMemberJoinConflict(testCase.ownerFamily, testCase.memberRecord)
|
||||
if testCase.wantCode == 0 {
|
||||
require.NoError(t, err)
|
||||
return
|
||||
}
|
||||
|
||||
require.Error(t, err)
|
||||
require.Equal(t, testCase.wantCode, extractFamilyJoinCode(err))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildRemovedSubscribeCacheMeta(t *testing.T) {
|
||||
removed := []modelUser.Subscribe{
|
||||
{Id: 1, SubscribeId: 10, Token: "member-token-1"},
|
||||
{Id: 2, SubscribeId: 11, Token: "member-token-2"},
|
||||
{Id: 3, SubscribeId: 0, Token: "member-token-3"},
|
||||
}
|
||||
|
||||
models, subscribeIDSet := buildRemovedSubscribeCacheMeta(removed)
|
||||
|
||||
require.Len(t, models, 3)
|
||||
require.Equal(t, int64(1), models[0].Id)
|
||||
require.Equal(t, "member-token-2", models[1].Token)
|
||||
require.Len(t, subscribeIDSet, 2)
|
||||
_, has10 := subscribeIDSet[10]
|
||||
_, has11 := subscribeIDSet[11]
|
||||
_, has0 := subscribeIDSet[0]
|
||||
require.True(t, has10)
|
||||
require.True(t, has11)
|
||||
require.False(t, has0)
|
||||
}
|
||||
@@ -1,105 +0,0 @@
|
||||
package user
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
modelUser "github.com/perfect-panel/server/internal/model/user"
|
||||
"github.com/perfect-panel/server/internal/types"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestAppendFamilyOwnerEmailIfNeeded(t *testing.T) {
|
||||
testCases := []struct {
|
||||
name string
|
||||
methods []types.UserAuthMethod
|
||||
familyJoined bool
|
||||
ownerEmailMethod *modelUser.AuthMethods
|
||||
wantMethodCount int
|
||||
wantEmailCount int
|
||||
wantFirstAuthType string
|
||||
wantFirstAuthValue string
|
||||
}{
|
||||
{
|
||||
name: "inject owner email when member has no email",
|
||||
methods: []types.UserAuthMethod{
|
||||
{AuthType: "device", AuthIdentifier: "dev-1", Verified: true},
|
||||
},
|
||||
familyJoined: true,
|
||||
ownerEmailMethod: &modelUser.AuthMethods{AuthType: "email", AuthIdentifier: "owner@example.com", Verified: true},
|
||||
wantMethodCount: 2,
|
||||
wantEmailCount: 1,
|
||||
wantFirstAuthType: "email",
|
||||
wantFirstAuthValue: "owner@example.com",
|
||||
},
|
||||
{
|
||||
name: "do not inject when member already has email",
|
||||
methods: []types.UserAuthMethod{
|
||||
{AuthType: "email", AuthIdentifier: "member@example.com", Verified: true},
|
||||
{AuthType: "device", AuthIdentifier: "dev-1", Verified: true},
|
||||
},
|
||||
familyJoined: true,
|
||||
ownerEmailMethod: &modelUser.AuthMethods{AuthType: "email", AuthIdentifier: "owner@example.com", Verified: true},
|
||||
wantMethodCount: 2,
|
||||
wantEmailCount: 1,
|
||||
wantFirstAuthType: "email",
|
||||
wantFirstAuthValue: "member@example.com",
|
||||
},
|
||||
{
|
||||
name: "do not inject when owner has no email",
|
||||
methods: []types.UserAuthMethod{
|
||||
{AuthType: "device", AuthIdentifier: "dev-1", Verified: true},
|
||||
},
|
||||
familyJoined: true,
|
||||
ownerEmailMethod: &modelUser.AuthMethods{AuthType: "email", AuthIdentifier: "", Verified: true},
|
||||
wantMethodCount: 1,
|
||||
wantEmailCount: 0,
|
||||
wantFirstAuthType: "device",
|
||||
},
|
||||
{
|
||||
name: "do not inject for non active family relationship",
|
||||
methods: []types.UserAuthMethod{
|
||||
{AuthType: "device", AuthIdentifier: "dev-1", Verified: true},
|
||||
},
|
||||
familyJoined: false,
|
||||
ownerEmailMethod: &modelUser.AuthMethods{AuthType: "email", AuthIdentifier: "owner@example.com", Verified: true},
|
||||
wantMethodCount: 1,
|
||||
wantEmailCount: 0,
|
||||
wantFirstAuthType: "device",
|
||||
},
|
||||
{
|
||||
name: "sort keeps injected email at first position",
|
||||
methods: []types.UserAuthMethod{
|
||||
{AuthType: "mobile", AuthIdentifier: "+1234567890", Verified: true},
|
||||
{AuthType: "device", AuthIdentifier: "dev-1", Verified: true},
|
||||
},
|
||||
familyJoined: true,
|
||||
ownerEmailMethod: &modelUser.AuthMethods{AuthType: "email", AuthIdentifier: "owner@example.com", Verified: true},
|
||||
wantMethodCount: 3,
|
||||
wantEmailCount: 1,
|
||||
wantFirstAuthType: "email",
|
||||
wantFirstAuthValue: "owner@example.com",
|
||||
},
|
||||
}
|
||||
|
||||
for _, testCase := range testCases {
|
||||
t.Run(testCase.name, func(t *testing.T) {
|
||||
finalMethods := appendFamilyOwnerEmailIfNeeded(testCase.methods, testCase.familyJoined, testCase.ownerEmailMethod)
|
||||
sortUserAuthMethodsByPriority(finalMethods)
|
||||
|
||||
require.Len(t, finalMethods, testCase.wantMethodCount)
|
||||
|
||||
emailCount := 0
|
||||
for _, method := range finalMethods {
|
||||
if method.AuthType == "email" {
|
||||
emailCount++
|
||||
}
|
||||
}
|
||||
require.Equal(t, testCase.wantEmailCount, emailCount)
|
||||
|
||||
require.Equal(t, testCase.wantFirstAuthType, finalMethods[0].AuthType)
|
||||
if testCase.wantFirstAuthValue != "" {
|
||||
require.Equal(t, testCase.wantFirstAuthValue, finalMethods[0].AuthIdentifier)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -1,25 +0,0 @@
|
||||
package user
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
commonLogic "github.com/perfect-panel/server/internal/logic/common"
|
||||
"github.com/perfect-panel/server/internal/types"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestFillUserSubscribeEntitlementFields(t *testing.T) {
|
||||
sub := &types.UserSubscribe{}
|
||||
entitlement := &commonLogic.EntitlementContext{
|
||||
EffectiveUserID: 2001,
|
||||
Source: commonLogic.EntitlementSourceFamilyOwner,
|
||||
OwnerUserID: 2001,
|
||||
ReadOnly: true,
|
||||
}
|
||||
|
||||
fillUserSubscribeEntitlementFields(sub, entitlement)
|
||||
|
||||
require.Equal(t, commonLogic.EntitlementSourceFamilyOwner, sub.EntitlementSource)
|
||||
require.Equal(t, int64(2001), sub.EntitlementOwnerUserId)
|
||||
require.True(t, sub.ReadOnly)
|
||||
}
|
||||
@@ -71,7 +71,7 @@ func (l *UnsubscribeLogic) Unsubscribe(req *types.UnsubscribeRequest) error {
|
||||
err = l.svcCtx.UserModel.Transaction(l.ctx, func(db *gorm.DB) error {
|
||||
// Find and update subscription status to cancelled (status = 4)
|
||||
userSub.Status = 4 // Set status to cancelled
|
||||
if err = l.svcCtx.UserModel.UpdateSubscribe(l.ctx, userSub); err != nil {
|
||||
if err = l.svcCtx.UserModel.UpdateSubscribe(l.ctx, userSub, db); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -148,7 +148,7 @@ func (l *UnsubscribeLogic) Unsubscribe(req *types.UnsubscribeRequest) error {
|
||||
|
||||
// Update user's regular balance and save changes to database
|
||||
u.Balance = balance
|
||||
return l.svcCtx.UserModel.Update(l.ctx, u)
|
||||
return l.svcCtx.UserModel.Update(l.ctx, u, db)
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
|
||||
Reference in New Issue
Block a user