354 lines
12 KiB
Go
354 lines
12 KiB
Go
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)
|
||
|
|
}
|