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) }