修复 UR#55 代码审查问题:事务隔离、错误包装、重复辅助函数

- resolvePackageTerms 接受 tx 参数并使用事务绑定的 Store,确保快照与写入同库
- auto_purchase.go 错误包装改用 pkgerrors.Wrap 而非标准库 errors
- shop_package_batch_allocation 去除 fmt 依赖改用 strconv 拼接
- shop_series_grant 提取 effectiveBase 局部变量避免重复调用
- 测试辅助函数统一迁移至 testutil.StringPointer,删除各文件本地重复定义
- 新增 testutil.NewRedisClient 和 testutil.StringPointer

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
2026-07-23 11:09:08 +09:00
parent f7e0f07692
commit 147f3eb775
7 changed files with 29 additions and 22 deletions

View File

@@ -32,8 +32,8 @@ func TestBatchAllocatePackagesHTTPRequiresExplicitExpiryBase(t *testing.T) {
expected *string expected *string
}{ }{
{name: "跟随默认", field: "null"}, {name: "跟随默认", field: "null"},
{name: "购买即生效", field: `"from_purchase"`, expected: stringPointer(constants.PackageExpiryBaseFromPurchase)}, {name: "购买即生效", field: `"from_purchase"`, expected: testutil.StringPointer(constants.PackageExpiryBaseFromPurchase)},
{name: "实名即生效", field: `"from_activation"`, expected: stringPointer(constants.PackageExpiryBaseFromActivation)}, {name: "实名即生效", field: `"from_activation"`, expected: testutil.StringPointer(constants.PackageExpiryBaseFromActivation)},
} }
for _, testCase := range testCases { for _, testCase := range testCases {
t.Run(testCase.name, func(t *testing.T) { t.Run(testCase.name, func(t *testing.T) {
@@ -204,8 +204,6 @@ func completeExpiryBaseTestUsage(packageID uint) *model.PackageUsage {
return &model.PackageUsage{OrderID: unique, OrderNo: "UR55-PATCH-USAGE", PackageID: packageID, UsageType: constants.AssetWalletResourceTypeIotCard, IotCardID: unique, DataLimitMB: 1, Status: constants.PackageUsageStatusPending, Priority: 1, PackageName: "UR55测试套餐", Generation: 1, ExpiryBaseSnapshot: constants.PackageExpiryBaseFromPurchase, CalendarTypeSnapshot: constants.PackageCalendarTypeByDay, DurationDaysSnapshot: 30} return &model.PackageUsage{OrderID: unique, OrderNo: "UR55-PATCH-USAGE", PackageID: packageID, UsageType: constants.AssetWalletResourceTypeIotCard, IotCardID: unique, DataLimitMB: 1, Status: constants.PackageUsageStatusPending, Priority: 1, PackageName: "UR55测试套餐", Generation: 1, ExpiryBaseSnapshot: constants.PackageExpiryBaseFromPurchase, CalendarTypeSnapshot: constants.PackageCalendarTypeByDay, DurationDaysSnapshot: 30}
} }
func stringPointer(value string) *string { return &value }
func nullableStringEqual(left, right *string) bool { func nullableStringEqual(left, right *string) bool {
return left == nil && right == nil || left != nil && right != nil && *left == *right return left == nil && right == nil || left != nil && right != nil && *left == *right
} }
@@ -218,8 +216,8 @@ func TestSeriesGrantCreateHTTPRequiresExplicitExpiryBase(t *testing.T) {
expected *string expected *string
}{ }{
{name: "跟随默认", field: "null"}, {name: "跟随默认", field: "null"},
{name: "购买即生效", field: `"from_purchase"`, expected: stringPointer(constants.PackageExpiryBaseFromPurchase)}, {name: "购买即生效", field: `"from_purchase"`, expected: testutil.StringPointer(constants.PackageExpiryBaseFromPurchase)},
{name: "实名即生效", field: `"from_activation"`, expected: stringPointer(constants.PackageExpiryBaseFromActivation)}, {name: "实名即生效", field: `"from_activation"`, expected: testutil.StringPointer(constants.PackageExpiryBaseFromActivation)},
} }
for _, testCase := range testCases { for _, testCase := range testCases {
t.Run(testCase.name, func(t *testing.T) { t.Run(testCase.name, func(t *testing.T) {

View File

@@ -2223,7 +2223,7 @@ func (s *Service) markOrderCreated(ctx context.Context, idempotencyKey string, o
// activateMainPackage 任务 8.2-8.4: 主套餐激活逻辑 // activateMainPackage 任务 8.2-8.4: 主套餐激活逻辑
func (s *Service) activateMainPackage(ctx context.Context, tx *gorm.DB, order *model.Order, pkg *model.Package, carrierType string, carrierID uint, now time.Time) error { func (s *Service) activateMainPackage(ctx context.Context, tx *gorm.DB, order *model.Order, pkg *model.Package, carrierType string, carrierID uint, now time.Time) error {
terms, err := s.resolvePackageTerms(ctx, pkg, order.SellerShopID) terms, err := s.resolvePackageTerms(ctx, tx, pkg, order.SellerShopID)
if err != nil { if err != nil {
s.logger.Error("生成套餐计时快照失败", zap.Uint("package_id", pkg.ID), zap.Error(err)) s.logger.Error("生成套餐计时快照失败", zap.Uint("package_id", pkg.ID), zap.Error(err))
return err return err
@@ -2366,7 +2366,7 @@ func (s *Service) activateMainPackage(ctx context.Context, tx *gorm.DB, order *m
// activateAddonPackage 任务 8.5-8.7: 加油包激活逻辑 // activateAddonPackage 任务 8.5-8.7: 加油包激活逻辑
func (s *Service) activateAddonPackage(ctx context.Context, tx *gorm.DB, order *model.Order, pkg *model.Package, carrierType string, carrierID uint, now time.Time) error { func (s *Service) activateAddonPackage(ctx context.Context, tx *gorm.DB, order *model.Order, pkg *model.Package, carrierType string, carrierID uint, now time.Time) error {
terms, err := s.resolvePackageTerms(ctx, pkg, order.SellerShopID) terms, err := s.resolvePackageTerms(ctx, tx, pkg, order.SellerShopID)
if err != nil { if err != nil {
s.logger.Error("生成加油包计时快照失败", zap.Uint("package_id", pkg.ID), zap.Error(err)) s.logger.Error("生成加油包计时快照失败", zap.Uint("package_id", pkg.ID), zap.Error(err))
return err return err
@@ -2446,10 +2446,13 @@ func (s *Service) activateAddonPackage(ctx context.Context, tx *gorm.DB, order *
return nil return nil
} }
func (s *Service) resolvePackageTerms(ctx context.Context, pkg *model.Package, sellerShopID *uint) (packagedomain.TermsSnapshot, error) { // resolvePackageTerms 在给定事务内查询分配覆盖并解析计时条款快照。
// 必须在事务内调用,确保分配读取与 PackageUsage 写入在同一连接,避免快照与提交值不一致。
func (s *Service) resolvePackageTerms(ctx context.Context, tx *gorm.DB, pkg *model.Package, sellerShopID *uint) (packagedomain.TermsSnapshot, error) {
var allocation *model.ShopPackageAllocation var allocation *model.ShopPackageAllocation
if sellerShopID != nil && *sellerShopID > 0 { if sellerShopID != nil && *sellerShopID > 0 {
found, err := s.shopPackageAllocationStore.GetByShopAndPackageForSystem(ctx, *sellerShopID, pkg.ID) store := postgres.NewShopPackageAllocationStore(tx)
found, err := store.GetByShopAndPackageForSystem(ctx, *sellerShopID, pkg.ID)
if err != nil && err != gorm.ErrRecordNotFound { if err != nil && err != gorm.ErrRecordNotFound {
return packagedomain.TermsSnapshot{}, errors.Wrap(errors.CodeDatabaseError, err, "查询套餐分配失败") return packagedomain.TermsSnapshot{}, errors.Wrap(errors.CodeDatabaseError, err, "查询套餐分配失败")
} }

View File

@@ -2,7 +2,7 @@ package shop_package_batch_allocation
import ( import (
"context" "context"
"fmt" "strconv"
"github.com/break/junhong_cmp_fiber/internal/model" "github.com/break/junhong_cmp_fiber/internal/model"
"github.com/break/junhong_cmp_fiber/internal/model/dto" "github.com/break/junhong_cmp_fiber/internal/model/dto"
@@ -94,7 +94,7 @@ func (s *Service) logExpiryBaseAudit(ctx context.Context, allocation *model.Shop
s.auditService.LogOperation(ctx, &model.AccountOperationLog{ s.auditService.LogOperation(ctx, &model.AccountOperationLog{
OperatorID: middleware.GetUserIDFromContext(ctx), OperatorType: middleware.GetUserTypeFromContext(ctx), OperatorID: middleware.GetUserIDFromContext(ctx), OperatorType: middleware.GetUserTypeFromContext(ctx),
OperatorName: middleware.GetUsernameFromContext(ctx), OperationType: "update_package_expiry_base", OperatorName: middleware.GetUsernameFromContext(ctx), OperationType: "update_package_expiry_base",
OperationDesc: fmt.Sprintf("修改套餐分配生效条件覆盖: %d", allocation.ID), OperationDesc: "修改套餐分配生效条件覆盖: " + strconv.FormatUint(uint64(allocation.ID), 10),
BeforeData: model.JSONB{"allocation_id": allocation.ID, "expiry_base_override": before}, BeforeData: model.JSONB{"allocation_id": allocation.ID, "expiry_base_override": before},
AfterData: model.JSONB{"allocation_id": allocation.ID, "expiry_base_override": after}, AfterData: model.JSONB{"allocation_id": allocation.ID, "expiry_base_override": after},
RequestID: middleware.GetRequestIDFromContext(ctx), IPAddress: middleware.GetIPFromContext(ctx), UserAgent: middleware.GetUserAgentFromContext(ctx), RequestID: middleware.GetRequestIDFromContext(ctx), IPAddress: middleware.GetIPFromContext(ctx), UserAgent: middleware.GetUserAgentFromContext(ctx),

View File

@@ -185,6 +185,7 @@ func (s *Service) buildGrantResponse(ctx context.Context, allocation *model.Shop
if pkg.IsGift { if pkg.IsGift {
continue continue
} }
effectiveBase := packagepkg.EffectiveExpiryBase(pkg, pa)
packages = append(packages, dto.ShopSeriesGrantPackageItem{ packages = append(packages, dto.ShopSeriesGrantPackageItem{
PackageID: pa.PackageID, PackageID: pa.PackageID,
PackageName: pkg.PackageName, PackageName: pkg.PackageName,
@@ -196,8 +197,8 @@ func (s *Service) buildGrantResponse(ctx context.Context, allocation *model.Shop
DefaultExpiryBaseName: packagepkg.ExpiryBaseName(pkg.ExpiryBase), DefaultExpiryBaseName: packagepkg.ExpiryBaseName(pkg.ExpiryBase),
ExpiryBaseOverride: pa.ExpiryBaseOverride, ExpiryBaseOverride: pa.ExpiryBaseOverride,
ExpiryBaseOverrideName: packagepkg.ExpiryBaseOverrideName(pa.ExpiryBaseOverride), ExpiryBaseOverrideName: packagepkg.ExpiryBaseOverrideName(pa.ExpiryBaseOverride),
EffectiveExpiryBase: packagepkg.EffectiveExpiryBase(pkg, pa), EffectiveExpiryBase: effectiveBase,
EffectiveExpiryBaseName: packagepkg.ExpiryBaseName(packagepkg.EffectiveExpiryBase(pkg, pa)), EffectiveExpiryBaseName: packagepkg.ExpiryBaseName(effectiveBase),
}) })
} }
resp.Packages = packages resp.Packages = packages

View File

@@ -498,7 +498,7 @@ func (h *AutoPurchaseHandler) activateMainPackage(
carrierID uint, carrierID uint,
now time.Time, now time.Time,
) error { ) error {
terms, err := h.resolvePackageTerms(ctx, pkg, order.SellerShopID) terms, err := h.resolvePackageTerms(ctx, tx, pkg, order.SellerShopID)
if err != nil { if err != nil {
h.logger.Error("自动购包生成套餐计时快照失败", zap.Uint("package_id", pkg.ID), zap.Error(err)) h.logger.Error("自动购包生成套餐计时快照失败", zap.Uint("package_id", pkg.ID), zap.Error(err))
return err return err
@@ -642,7 +642,7 @@ func (h *AutoPurchaseHandler) activateAddonPackage(
carrierID uint, carrierID uint,
now time.Time, now time.Time,
) error { ) error {
terms, err := h.resolvePackageTerms(ctx, pkg, order.SellerShopID) terms, err := h.resolvePackageTerms(ctx, tx, pkg, order.SellerShopID)
if err != nil { if err != nil {
h.logger.Error("自动购包生成加油包计时快照失败", zap.Uint("package_id", pkg.ID), zap.Error(err)) h.logger.Error("自动购包生成加油包计时快照失败", zap.Uint("package_id", pkg.ID), zap.Error(err))
return err return err
@@ -703,12 +703,15 @@ func (h *AutoPurchaseHandler) activateAddonPackage(
return tx.Create(usage).Error return tx.Create(usage).Error
} }
func (h *AutoPurchaseHandler) resolvePackageTerms(ctx context.Context, pkg *model.Package, sellerShopID *uint) (packagedomain.TermsSnapshot, error) { // resolvePackageTerms 在给定事务内查询分配覆盖并解析计时条款快照。
// 必须在事务内调用,确保分配读取与 PackageUsage 写入在同一连接,避免快照与提交值不一致。
func (h *AutoPurchaseHandler) resolvePackageTerms(ctx context.Context, tx *gorm.DB, pkg *model.Package, sellerShopID *uint) (packagedomain.TermsSnapshot, error) {
var allocation *model.ShopPackageAllocation var allocation *model.ShopPackageAllocation
if sellerShopID != nil && *sellerShopID > 0 { if sellerShopID != nil && *sellerShopID > 0 {
found, err := h.shopPackageAllocationStore.GetByShopAndPackageForSystem(ctx, *sellerShopID, pkg.ID) store := postgres.NewShopPackageAllocationStore(tx)
found, err := store.GetByShopAndPackageForSystem(ctx, *sellerShopID, pkg.ID)
if err != nil && err != gorm.ErrRecordNotFound { if err != nil && err != gorm.ErrRecordNotFound {
return packagedomain.TermsSnapshot{}, err return packagedomain.TermsSnapshot{}, pkgerrors.Wrap(pkgerrors.CodeDatabaseError, err, "查询套餐分配失败")
} }
if err == nil { if err == nil {
allocation = found allocation = found

View File

@@ -26,8 +26,8 @@ func TestAutoPurchasePersistsTermsSnapshotsAndRealnameDecision(t *testing.T) {
expectedPending bool expectedPending bool
}{ }{
{name: "跟随默认等待实名", defaultBase: constants.PackageExpiryBaseFromActivation, realnameStatus: constants.RealNameStatusNotVerified, expectedBase: constants.PackageExpiryBaseFromActivation, expectedStatus: constants.PackageUsageStatusPending, expectedPending: true}, {name: "跟随默认等待实名", defaultBase: constants.PackageExpiryBaseFromActivation, realnameStatus: constants.RealNameStatusNotVerified, expectedBase: constants.PackageExpiryBaseFromActivation, expectedStatus: constants.PackageUsageStatusPending, expectedPending: true},
{name: "覆盖购买即生效", defaultBase: constants.PackageExpiryBaseFromActivation, override: taskStringPointer(constants.PackageExpiryBaseFromPurchase), realnameStatus: constants.RealNameStatusNotVerified, expectedBase: constants.PackageExpiryBaseFromPurchase, expectedStatus: constants.PackageUsageStatusActive}, {name: "覆盖购买即生效", defaultBase: constants.PackageExpiryBaseFromActivation, override: testutil.StringPointer(constants.PackageExpiryBaseFromPurchase), realnameStatus: constants.RealNameStatusNotVerified, expectedBase: constants.PackageExpiryBaseFromPurchase, expectedStatus: constants.PackageUsageStatusActive},
{name: "覆盖实名但已实名", defaultBase: constants.PackageExpiryBaseFromPurchase, override: taskStringPointer(constants.PackageExpiryBaseFromActivation), realnameStatus: constants.RealNameStatusVerified, expectedBase: constants.PackageExpiryBaseFromActivation, expectedStatus: constants.PackageUsageStatusActive}, {name: "覆盖实名但已实名", defaultBase: constants.PackageExpiryBaseFromPurchase, override: testutil.StringPointer(constants.PackageExpiryBaseFromActivation), realnameStatus: constants.RealNameStatusVerified, expectedBase: constants.PackageExpiryBaseFromActivation, expectedStatus: constants.PackageUsageStatusActive},
} }
for _, testCase := range testCases { for _, testCase := range testCases {
t.Run(testCase.name, func(t *testing.T) { t.Run(testCase.name, func(t *testing.T) {
@@ -166,4 +166,3 @@ func newAutoPurchaseOrder(cardID uint, sellerShopID *uint) *model.Order {
return &model.Order{Model: gorm.Model{ID: orderID}, OrderNo: "UR55-AUTO-" + strconv.FormatUint(uint64(orderID), 10), OrderType: model.OrderTypeSingleCard, IotCardID: &cardID, SellerShopID: sellerShopID, TotalAmount: 100, Generation: 1} return &model.Order{Model: gorm.Model{ID: orderID}, OrderNo: "UR55-AUTO-" + strconv.FormatUint(uint64(orderID), 10), OrderType: model.OrderTypeSingleCard, IotCardID: &cardID, SellerShopID: sellerShopID, TotalAmount: 100, Generation: 1}
} }
func taskStringPointer(value string) *string { return &value }

View File

@@ -61,3 +61,6 @@ func NewRedisClient(t *testing.T) *redis.Client {
t.Cleanup(func() { _ = client.Close() }) t.Cleanup(func() { _ = client.Close() })
return client return client
} }
// StringPointer 返回字符串值的指针,供测试用例构造 nullable string 参数。
func StringPointer(s string) *string { return &s }