feat(wallet): add wallet entities and repository
Entities for the four wallet tables and a WalletRepository that the wallet processor will build on (PC-103). Every method goes through the caller's transaction, and writes and locks refuse to run without one: outside a transaction a lock is released as soon as it is taken and a balance could move without its ledger row. - LockWallet creates the wallet on first use, taking the organization from the customer, then locks it with SELECT ... FOR UPDATE. - LockWallets always locks in customer_id order so opposite transfers cannot deadlock. - AddBalance and ConsumeLot are conditional updates that return an error when they would overdraw, instead of tripping the CHECK constraint. - ListActiveLots returns unexpired lots with balance in K9 spending order. The tests need a real Postgres and run only when TEST_DATABASE_URL points at a migrated database. Both the lock and the lock ordering were checked by removing them and watching the tests fail (lost update, deadlock detected). Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 5.5
parent
b107f4ef04
commit
fc5eecb68a
@@ -0,0 +1,261 @@
|
||||
package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"sort"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/clause"
|
||||
|
||||
"apskel-pos-be/internal/constants"
|
||||
"apskel-pos-be/internal/entities"
|
||||
)
|
||||
|
||||
var (
|
||||
// ErrWalletTxRequired is returned by every write and lock when the context carries
|
||||
// no transaction from TxManager. Outside a transaction a lock is released as soon as
|
||||
// it is taken, and a balance could move without its ledger row.
|
||||
ErrWalletTxRequired = errors.New("wallet: operation must run inside a transaction")
|
||||
// ErrWalletNotFound means the customer does not exist, so no wallet could be made.
|
||||
ErrWalletNotFound = errors.New("wallet: customer not found")
|
||||
// ErrWalletInsufficientBalance means a conditional update matched no row because
|
||||
// the balance would have gone negative.
|
||||
ErrWalletInsufficientBalance = errors.New("wallet: insufficient balance")
|
||||
// ErrWalletLotInsufficient means a lot had less remaining than was taken from it.
|
||||
ErrWalletLotInsufficient = errors.New("wallet: lot has insufficient remaining amount")
|
||||
)
|
||||
|
||||
// WalletRepository reads and writes the wallet tables (docs/prd-point-coin.md §7, §8).
|
||||
// Only the wallet processor should call its write methods: it is the one place that
|
||||
// keeps balances, ledger rows and lots in step.
|
||||
//
|
||||
// Unlike the gamification repositories, every method goes through DBFromContext so it
|
||||
// joins the caller's transaction. Writes and locks refuse to run without one.
|
||||
type WalletRepository interface {
|
||||
// LockWallet locks the customer's wallet row for the rest of the transaction,
|
||||
// creating the row first if the customer has none yet.
|
||||
LockWallet(ctx context.Context, customerID uuid.UUID) (*entities.CustomerWallet, error)
|
||||
// LockWallets locks two wallets, always in customer_id order so that two transfers
|
||||
// in opposite directions cannot deadlock. The results come back in argument order.
|
||||
LockWallets(ctx context.Context, a, b uuid.UUID) (*entities.CustomerWallet, *entities.CustomerWallet, error)
|
||||
// AddBalance moves one balance by delta and returns the new balance. A debit that
|
||||
// would make it negative changes nothing and returns ErrWalletInsufficientBalance.
|
||||
AddBalance(ctx context.Context, customerID uuid.UUID, currency string, delta int64) (int64, error)
|
||||
GetWallet(ctx context.Context, customerID uuid.UUID) (*entities.CustomerWallet, error)
|
||||
|
||||
CreateTransaction(ctx context.Context, walletTx *entities.WalletTransaction) error
|
||||
// GetTransactionByIdempotencyKey returns nil, nil when no row has the key.
|
||||
GetTransactionByIdempotencyKey(ctx context.Context, key string) (*entities.WalletTransaction, error)
|
||||
|
||||
CreateLot(ctx context.Context, lot *entities.WalletLot) error
|
||||
// ListActiveLots returns the lots that still have balance and have not expired at
|
||||
// asOf, in the order they are spent (K9): soonest expiry first, lots without an
|
||||
// expiry last, oldest first within the same expiry.
|
||||
ListActiveLots(ctx context.Context, customerID uuid.UUID, currency string, asOf time.Time) ([]entities.WalletLot, error)
|
||||
// ConsumeLot takes amount from a lot's remaining amount. Taking more than remains
|
||||
// changes nothing and returns ErrWalletLotInsufficient.
|
||||
ConsumeLot(ctx context.Context, lotID uuid.UUID, amount int64) error
|
||||
|
||||
CreateAllocations(ctx context.Context, allocations []entities.WalletLotAllocation) error
|
||||
ListAllocationsByTransaction(ctx context.Context, transactionID uuid.UUID) ([]entities.WalletLotAllocation, error)
|
||||
}
|
||||
|
||||
type walletRepository struct {
|
||||
db *gorm.DB
|
||||
}
|
||||
|
||||
func NewWalletRepository(db *gorm.DB) WalletRepository {
|
||||
return &walletRepository{db: db}
|
||||
}
|
||||
|
||||
// txDB returns the caller's transaction, or ErrWalletTxRequired if there is none.
|
||||
func (r *walletRepository) txDB(ctx context.Context) (*gorm.DB, error) {
|
||||
if tx, ok := ctx.Value(txKey).(*gorm.DB); ok && tx != nil {
|
||||
return tx.WithContext(ctx), nil
|
||||
}
|
||||
return nil, ErrWalletTxRequired
|
||||
}
|
||||
|
||||
func (r *walletRepository) LockWallet(ctx context.Context, customerID uuid.UUID) (*entities.CustomerWallet, error) {
|
||||
db, err := r.txDB(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// The wallet takes its organization from the customer, so the two cannot disagree.
|
||||
// ON CONFLICT covers two first operations racing to create the same wallet.
|
||||
err = db.Exec(`INSERT INTO customer_wallets (customer_id, organization_id)
|
||||
SELECT id, organization_id FROM customers WHERE id = ?
|
||||
ON CONFLICT (customer_id) DO NOTHING`, customerID).Error
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create customer wallet: %w", err)
|
||||
}
|
||||
|
||||
var wallet entities.CustomerWallet
|
||||
err = db.Clauses(clause.Locking{Strength: "UPDATE"}).
|
||||
Where("customer_id = ?", customerID).
|
||||
First(&wallet).Error
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, ErrWalletNotFound
|
||||
}
|
||||
return nil, fmt.Errorf("failed to lock customer wallet: %w", err)
|
||||
}
|
||||
return &wallet, nil
|
||||
}
|
||||
|
||||
func (r *walletRepository) LockWallets(ctx context.Context, a, b uuid.UUID) (*entities.CustomerWallet, *entities.CustomerWallet, error) {
|
||||
if a == b {
|
||||
return nil, nil, errors.New("wallet: cannot lock the same wallet twice")
|
||||
}
|
||||
|
||||
ids := []uuid.UUID{a, b}
|
||||
sort.Slice(ids, func(i, j int) bool { return ids[i].String() < ids[j].String() })
|
||||
|
||||
locked := make(map[uuid.UUID]*entities.CustomerWallet, 2)
|
||||
for _, id := range ids {
|
||||
wallet, err := r.LockWallet(ctx, id)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
locked[id] = wallet
|
||||
}
|
||||
return locked[a], locked[b], nil
|
||||
}
|
||||
|
||||
func (r *walletRepository) AddBalance(ctx context.Context, customerID uuid.UUID, currency string, delta int64) (int64, error) {
|
||||
db, err := r.txDB(ctx)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
|
||||
var column string
|
||||
switch currency {
|
||||
case constants.WalletCurrencyPoint:
|
||||
column = "point_balance"
|
||||
case constants.WalletCurrencyCoin:
|
||||
column = "coin_balance"
|
||||
default:
|
||||
return 0, fmt.Errorf("wallet: unknown currency %q", currency)
|
||||
}
|
||||
|
||||
// The WHERE clause makes an overdraft match no row instead of tripping the CHECK,
|
||||
// so the caller gets a clean error and the transaction stays usable.
|
||||
var balances []int64
|
||||
err = db.Raw(`UPDATE customer_wallets
|
||||
SET `+column+` = `+column+` + ?, updated_at = NOW()
|
||||
WHERE customer_id = ? AND `+column+` + ? >= 0
|
||||
RETURNING `+column, delta, customerID, delta).
|
||||
Scan(&balances).Error
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("failed to update wallet balance: %w", err)
|
||||
}
|
||||
if len(balances) == 0 {
|
||||
if delta >= 0 {
|
||||
return 0, ErrWalletNotFound
|
||||
}
|
||||
return 0, ErrWalletInsufficientBalance
|
||||
}
|
||||
return balances[0], nil
|
||||
}
|
||||
|
||||
func (r *walletRepository) GetWallet(ctx context.Context, customerID uuid.UUID) (*entities.CustomerWallet, error) {
|
||||
var wallet entities.CustomerWallet
|
||||
err := DBFromContext(ctx, r.db).WithContext(ctx).
|
||||
Where("customer_id = ?", customerID).
|
||||
First(&wallet).Error
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &wallet, nil
|
||||
}
|
||||
|
||||
func (r *walletRepository) CreateTransaction(ctx context.Context, walletTx *entities.WalletTransaction) error {
|
||||
db, err := r.txDB(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return db.Create(walletTx).Error
|
||||
}
|
||||
|
||||
func (r *walletRepository) GetTransactionByIdempotencyKey(ctx context.Context, key string) (*entities.WalletTransaction, error) {
|
||||
var walletTx entities.WalletTransaction
|
||||
err := DBFromContext(ctx, r.db).WithContext(ctx).
|
||||
Where("idempotency_key = ?", key).
|
||||
First(&walletTx).Error
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, nil
|
||||
}
|
||||
return nil, fmt.Errorf("failed to get wallet transaction by idempotency key: %w", err)
|
||||
}
|
||||
return &walletTx, nil
|
||||
}
|
||||
|
||||
func (r *walletRepository) CreateLot(ctx context.Context, lot *entities.WalletLot) error {
|
||||
db, err := r.txDB(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return db.Create(lot).Error
|
||||
}
|
||||
|
||||
func (r *walletRepository) ListActiveLots(ctx context.Context, customerID uuid.UUID, currency string, asOf time.Time) ([]entities.WalletLot, error) {
|
||||
var lots []entities.WalletLot
|
||||
// Filter and order match idx_wallet_lots_consume.
|
||||
err := DBFromContext(ctx, r.db).WithContext(ctx).
|
||||
Where("customer_id = ? AND currency = ? AND remaining_amount > 0", customerID, currency).
|
||||
Where("(expires_at IS NULL OR expires_at > ?)", asOf).
|
||||
Order("expires_at NULLS LAST, created_at, id").
|
||||
Find(&lots).Error
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to list active wallet lots: %w", err)
|
||||
}
|
||||
return lots, nil
|
||||
}
|
||||
|
||||
func (r *walletRepository) ConsumeLot(ctx context.Context, lotID uuid.UUID, amount int64) error {
|
||||
if amount <= 0 {
|
||||
return fmt.Errorf("wallet: lot consumption must be positive, got %d", amount)
|
||||
}
|
||||
db, err := r.txDB(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
result := db.Exec(`UPDATE wallet_lots SET remaining_amount = remaining_amount - ?
|
||||
WHERE id = ? AND remaining_amount >= ?`, amount, lotID, amount)
|
||||
if result.Error != nil {
|
||||
return fmt.Errorf("failed to consume wallet lot: %w", result.Error)
|
||||
}
|
||||
if result.RowsAffected == 0 {
|
||||
return ErrWalletLotInsufficient
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *walletRepository) CreateAllocations(ctx context.Context, allocations []entities.WalletLotAllocation) error {
|
||||
if len(allocations) == 0 {
|
||||
return nil
|
||||
}
|
||||
db, err := r.txDB(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return db.Create(&allocations).Error
|
||||
}
|
||||
|
||||
func (r *walletRepository) ListAllocationsByTransaction(ctx context.Context, transactionID uuid.UUID) ([]entities.WalletLotAllocation, error) {
|
||||
var allocations []entities.WalletLotAllocation
|
||||
err := DBFromContext(ctx, r.db).WithContext(ctx).
|
||||
Where("transaction_id = ?", transactionID).
|
||||
Find(&allocations).Error
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to list wallet lot allocations: %w", err)
|
||||
}
|
||||
return allocations, nil
|
||||
}
|
||||
@@ -0,0 +1,353 @@
|
||||
package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"os"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/driver/postgres"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/logger"
|
||||
|
||||
"apskel-pos-be/internal/constants"
|
||||
"apskel-pos-be/internal/entities"
|
||||
)
|
||||
|
||||
// These tests need a real Postgres, because what they check (row locks and
|
||||
// conditional updates) only exists there. Point TEST_DATABASE_URL at a database with
|
||||
// all migrations applied, e.g.
|
||||
//
|
||||
// TEST_DATABASE_URL=postgres://user:pass@localhost:5432/pos_test?sslmode=disable go test ./internal/repository/ -run Wallet
|
||||
//
|
||||
// Each test creates its own organization and customers and removes them afterwards.
|
||||
func walletTestDB(t *testing.T) *gorm.DB {
|
||||
t.Helper()
|
||||
dsn := os.Getenv("TEST_DATABASE_URL")
|
||||
if dsn == "" {
|
||||
t.Skip("TEST_DATABASE_URL not set")
|
||||
}
|
||||
db, err := gorm.Open(postgres.Open(dsn), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)})
|
||||
require.NoError(t, err)
|
||||
return db
|
||||
}
|
||||
|
||||
type walletFixture struct {
|
||||
db *gorm.DB
|
||||
repo WalletRepository
|
||||
txm *TxManager
|
||||
orgID uuid.UUID
|
||||
customers []uuid.UUID
|
||||
}
|
||||
|
||||
func newWalletFixture(t *testing.T, customerCount int) *walletFixture {
|
||||
t.Helper()
|
||||
db := walletTestDB(t)
|
||||
f := &walletFixture{db: db, repo: NewWalletRepository(db), txm: NewTxManager(db), orgID: uuid.New()}
|
||||
|
||||
require.NoError(t, db.Exec(`INSERT INTO organizations (id, name, plan_type) VALUES (?, 'wallet test', 'basic')`, f.orgID).Error)
|
||||
for i := 0; i < customerCount; i++ {
|
||||
id := uuid.New()
|
||||
require.NoError(t, db.Exec(`INSERT INTO customers (id, organization_id, name) VALUES (?, ?, 'wallet test')`, id, f.orgID).Error)
|
||||
f.customers = append(f.customers, id)
|
||||
}
|
||||
|
||||
t.Cleanup(func() {
|
||||
for _, q := range []string{
|
||||
`DELETE FROM wallet_lot_allocations WHERE lot_id IN (SELECT id FROM wallet_lots WHERE customer_id IN ?)`,
|
||||
`DELETE FROM wallet_lots WHERE customer_id IN ?`,
|
||||
`DELETE FROM wallet_transactions WHERE customer_id IN ?`,
|
||||
`DELETE FROM customer_wallets WHERE customer_id IN ?`,
|
||||
`DELETE FROM customers WHERE id IN ?`,
|
||||
} {
|
||||
db.Exec(q, f.customers)
|
||||
}
|
||||
db.Exec(`DELETE FROM organizations WHERE id = ?`, f.orgID)
|
||||
})
|
||||
return f
|
||||
}
|
||||
|
||||
// inTx runs fn in a transaction and fails the test on error.
|
||||
func (f *walletFixture) inTx(t *testing.T, fn func(ctx context.Context) error) {
|
||||
t.Helper()
|
||||
require.NoError(t, f.txm.WithTransaction(context.Background(), fn))
|
||||
}
|
||||
|
||||
// credit writes a ledger row and a lot and moves the balance, the minimum the
|
||||
// database accepts for a credit.
|
||||
func (f *walletFixture) credit(t *testing.T, ctx context.Context, customerID uuid.UUID, amount int64, expiresAt *time.Time) *entities.WalletLot {
|
||||
t.Helper()
|
||||
balance, err := f.repo.AddBalance(ctx, customerID, constants.WalletCurrencyPoint, amount)
|
||||
require.NoError(t, err)
|
||||
walletTx := &entities.WalletTransaction{
|
||||
OrganizationID: f.orgID,
|
||||
CustomerID: customerID,
|
||||
Currency: constants.WalletCurrencyPoint,
|
||||
Type: constants.WalletTxTypeMigration,
|
||||
Amount: amount,
|
||||
BalanceAfter: balance,
|
||||
ReferenceType: constants.WalletRefTypeLegacyPoints,
|
||||
ReferenceID: uuid.New(),
|
||||
Description: "test",
|
||||
}
|
||||
require.NoError(t, f.repo.CreateTransaction(ctx, walletTx))
|
||||
lot := &entities.WalletLot{
|
||||
OrganizationID: f.orgID,
|
||||
CustomerID: customerID,
|
||||
Currency: constants.WalletCurrencyPoint,
|
||||
SourceTransactionID: walletTx.ID,
|
||||
OriginalAmount: amount,
|
||||
RemainingAmount: amount,
|
||||
ExpiresAt: expiresAt,
|
||||
}
|
||||
require.NoError(t, f.repo.CreateLot(ctx, lot))
|
||||
return lot
|
||||
}
|
||||
|
||||
func TestWalletRepository_WritesRequireTransaction(t *testing.T) {
|
||||
f := newWalletFixture(t, 1)
|
||||
ctx := context.Background()
|
||||
|
||||
_, err := f.repo.LockWallet(ctx, f.customers[0])
|
||||
assert.ErrorIs(t, err, ErrWalletTxRequired)
|
||||
_, err = f.repo.AddBalance(ctx, f.customers[0], constants.WalletCurrencyPoint, 10)
|
||||
assert.ErrorIs(t, err, ErrWalletTxRequired)
|
||||
assert.ErrorIs(t, f.repo.ConsumeLot(ctx, uuid.New(), 1), ErrWalletTxRequired)
|
||||
assert.ErrorIs(t, f.repo.CreateTransaction(ctx, &entities.WalletTransaction{}), ErrWalletTxRequired)
|
||||
assert.ErrorIs(t, f.repo.CreateLot(ctx, &entities.WalletLot{}), ErrWalletTxRequired)
|
||||
}
|
||||
|
||||
func TestWalletRepository_LockWalletCreatesWallet(t *testing.T) {
|
||||
f := newWalletFixture(t, 1)
|
||||
|
||||
f.inTx(t, func(ctx context.Context) error {
|
||||
wallet, err := f.repo.LockWallet(ctx, f.customers[0])
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, f.orgID, wallet.OrganizationID, "organization comes from the customer")
|
||||
assert.Zero(t, wallet.PointBalance)
|
||||
assert.Zero(t, wallet.CoinBalance)
|
||||
|
||||
// Locking again in the same transaction finds the same row.
|
||||
again, err := f.repo.LockWallet(ctx, f.customers[0])
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, wallet.CustomerID, again.CustomerID)
|
||||
return nil
|
||||
})
|
||||
|
||||
f.inTx(t, func(ctx context.Context) error {
|
||||
_, err := f.repo.LockWallet(ctx, uuid.New())
|
||||
assert.ErrorIs(t, err, ErrWalletNotFound)
|
||||
return nil
|
||||
})
|
||||
}
|
||||
|
||||
func TestWalletRepository_AddBalanceRejectsOverdraft(t *testing.T) {
|
||||
f := newWalletFixture(t, 1)
|
||||
customerID := f.customers[0]
|
||||
|
||||
f.inTx(t, func(ctx context.Context) error {
|
||||
_, err := f.repo.LockWallet(ctx, customerID)
|
||||
require.NoError(t, err)
|
||||
|
||||
balance, err := f.repo.AddBalance(ctx, customerID, constants.WalletCurrencyPoint, 5)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, int64(5), balance)
|
||||
|
||||
_, err = f.repo.AddBalance(ctx, customerID, constants.WalletCurrencyPoint, -6)
|
||||
assert.ErrorIs(t, err, ErrWalletInsufficientBalance)
|
||||
|
||||
// Coin is a separate balance: point balance does not cover it.
|
||||
_, err = f.repo.AddBalance(ctx, customerID, constants.WalletCurrencyCoin, -1)
|
||||
assert.ErrorIs(t, err, ErrWalletInsufficientBalance)
|
||||
|
||||
// The failed update left the transaction usable and the balance untouched.
|
||||
balance, err = f.repo.AddBalance(ctx, customerID, constants.WalletCurrencyPoint, -5)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, int64(0), balance)
|
||||
return nil
|
||||
})
|
||||
|
||||
wallet, err := f.repo.GetWallet(context.Background(), customerID)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, int64(0), wallet.PointBalance)
|
||||
assert.Equal(t, int64(0), wallet.CoinBalance)
|
||||
}
|
||||
|
||||
func TestWalletRepository_AddBalanceWithoutWallet(t *testing.T) {
|
||||
f := newWalletFixture(t, 1)
|
||||
|
||||
f.inTx(t, func(ctx context.Context) error {
|
||||
_, err := f.repo.AddBalance(ctx, f.customers[0], constants.WalletCurrencyPoint, 5)
|
||||
assert.ErrorIs(t, err, ErrWalletNotFound)
|
||||
_, err = f.repo.AddBalance(ctx, f.customers[0], "GOLD", 5)
|
||||
assert.Error(t, err)
|
||||
return nil
|
||||
})
|
||||
}
|
||||
|
||||
func TestWalletRepository_ConsumeLotRejectsOverdraw(t *testing.T) {
|
||||
f := newWalletFixture(t, 1)
|
||||
customerID := f.customers[0]
|
||||
|
||||
f.inTx(t, func(ctx context.Context) error {
|
||||
_, err := f.repo.LockWallet(ctx, customerID)
|
||||
require.NoError(t, err)
|
||||
lot := f.credit(t, ctx, customerID, 10, nil)
|
||||
|
||||
require.NoError(t, f.repo.ConsumeLot(ctx, lot.ID, 4))
|
||||
assert.ErrorIs(t, f.repo.ConsumeLot(ctx, lot.ID, 7), ErrWalletLotInsufficient)
|
||||
require.NoError(t, f.repo.ConsumeLot(ctx, lot.ID, 6))
|
||||
assert.ErrorIs(t, f.repo.ConsumeLot(ctx, lot.ID, 1), ErrWalletLotInsufficient)
|
||||
assert.Error(t, f.repo.ConsumeLot(ctx, lot.ID, 0))
|
||||
return nil
|
||||
})
|
||||
}
|
||||
|
||||
// Two goroutines lock the same wallet and do a read-modify-write with a pause in
|
||||
// between. Without the lock both would read 0 and the result would be 1.
|
||||
func TestWalletRepository_LockWalletSerializes(t *testing.T) {
|
||||
f := newWalletFixture(t, 1)
|
||||
customerID := f.customers[0]
|
||||
|
||||
// Create the wallet up front. Otherwise the second goroutine's INSERT ... ON
|
||||
// CONFLICT waits on the first one's uncommitted insert, which serializes them
|
||||
// even without FOR UPDATE and the test would prove nothing about the lock.
|
||||
f.inTx(t, func(ctx context.Context) error {
|
||||
_, err := f.repo.LockWallet(ctx, customerID)
|
||||
return err
|
||||
})
|
||||
|
||||
type window struct{ locked, released time.Time }
|
||||
windows := make([]window, 2)
|
||||
var wg sync.WaitGroup
|
||||
errs := make(chan error, 2)
|
||||
|
||||
for i := 0; i < 2; i++ {
|
||||
wg.Add(1)
|
||||
go func(i int) {
|
||||
defer wg.Done()
|
||||
errs <- f.txm.WithTransaction(context.Background(), func(ctx context.Context) error {
|
||||
wallet, err := f.repo.LockWallet(ctx, customerID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
windows[i].locked = time.Now()
|
||||
time.Sleep(300 * time.Millisecond)
|
||||
db := DBFromContext(ctx, f.db)
|
||||
if err := db.Exec(`UPDATE customer_wallets SET point_balance = ? WHERE customer_id = ?`,
|
||||
wallet.PointBalance+1, customerID).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
windows[i].released = time.Now()
|
||||
return nil
|
||||
})
|
||||
}(i)
|
||||
}
|
||||
wg.Wait()
|
||||
close(errs)
|
||||
for err := range errs {
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
wallet, err := f.repo.GetWallet(context.Background(), customerID)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, int64(2), wallet.PointBalance, "second transaction must see the first one's write")
|
||||
|
||||
first, second := windows[0], windows[1]
|
||||
if second.locked.Before(first.locked) {
|
||||
first, second = second, first
|
||||
}
|
||||
assert.False(t, second.locked.Before(first.released), "second lock was taken while the first was held")
|
||||
}
|
||||
|
||||
// Transfers in opposite directions lock the same pair of wallets. Because LockWallets
|
||||
// always locks in customer_id order, they queue instead of deadlocking.
|
||||
func TestWalletRepository_LockWalletsOppositeOrderDoesNotDeadlock(t *testing.T) {
|
||||
f := newWalletFixture(t, 2)
|
||||
a, b := f.customers[0], f.customers[1]
|
||||
// Existing wallets, for the same reason as in LockWalletSerializes.
|
||||
f.inTx(t, func(ctx context.Context) error {
|
||||
_, _, err := f.repo.LockWallets(ctx, a, b)
|
||||
return err
|
||||
})
|
||||
|
||||
var wg sync.WaitGroup
|
||||
errs := make(chan error, 20)
|
||||
for i := 0; i < 10; i++ {
|
||||
for _, pair := range [][2]uuid.UUID{{a, b}, {b, a}} {
|
||||
wg.Add(1)
|
||||
go func(first, second uuid.UUID) {
|
||||
defer wg.Done()
|
||||
errs <- f.txm.WithTransaction(context.Background(), func(ctx context.Context) error {
|
||||
w1, w2, err := f.repo.LockWallets(ctx, first, second)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if w1.CustomerID != first || w2.CustomerID != second {
|
||||
return errors.New("wallets returned out of argument order")
|
||||
}
|
||||
time.Sleep(20 * time.Millisecond)
|
||||
return nil
|
||||
})
|
||||
}(pair[0], pair[1])
|
||||
}
|
||||
}
|
||||
wg.Wait()
|
||||
close(errs)
|
||||
for err := range errs {
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
f.inTx(t, func(ctx context.Context) error {
|
||||
_, _, err := f.repo.LockWallets(ctx, a, a)
|
||||
assert.Error(t, err)
|
||||
return nil
|
||||
})
|
||||
}
|
||||
|
||||
func TestWalletRepository_ListActiveLotsOrder(t *testing.T) {
|
||||
f := newWalletFixture(t, 1)
|
||||
customerID := f.customers[0]
|
||||
now := time.Now()
|
||||
at := func(d time.Duration) *time.Time { v := now.Add(d); return &v }
|
||||
|
||||
create := func(expiresAt *time.Time) *entities.WalletLot {
|
||||
var lot *entities.WalletLot
|
||||
f.inTx(t, func(ctx context.Context) error {
|
||||
_, err := f.repo.LockWallet(ctx, customerID)
|
||||
require.NoError(t, err)
|
||||
lot = f.credit(t, ctx, customerID, 10, expiresAt)
|
||||
return nil
|
||||
})
|
||||
return lot
|
||||
}
|
||||
neverOld := create(nil)
|
||||
late := create(at(48 * time.Hour))
|
||||
soon := create(at(time.Hour))
|
||||
neverNew := create(nil)
|
||||
expired := create(at(-time.Hour))
|
||||
empty := create(at(30 * time.Minute))
|
||||
f.inTx(t, func(ctx context.Context) error {
|
||||
return f.repo.ConsumeLot(ctx, empty.ID, 10)
|
||||
})
|
||||
|
||||
lots, err := f.repo.ListActiveLots(context.Background(), customerID, constants.WalletCurrencyPoint, now)
|
||||
require.NoError(t, err)
|
||||
|
||||
var got []uuid.UUID
|
||||
for _, lot := range lots {
|
||||
got = append(got, lot.ID)
|
||||
}
|
||||
assert.Equal(t, []uuid.UUID{soon.ID, late.ID, neverOld.ID, neverNew.ID}, got,
|
||||
"soonest expiry first, no expiry last and oldest first, expired and empty lots left out")
|
||||
assert.NotContains(t, got, expired.ID)
|
||||
|
||||
coinLots, err := f.repo.ListActiveLots(context.Background(), customerID, constants.WalletCurrencyCoin, now)
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, coinLots)
|
||||
}
|
||||
Reference in New Issue
Block a user