From fc5eecb68ace46d984f86a6565c986599820cf8c Mon Sep 17 00:00:00 2001 From: efrilm Date: Wed, 30 Sep 2026 01:01:03 +0700 Subject: [PATCH] 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 --- internal/constants/wallet.go | 44 +++ internal/entities/entities.go | 5 + internal/entities/wallet.go | 113 ++++++ internal/repository/wallet_repository.go | 261 +++++++++++++ internal/repository/wallet_repository_test.go | 353 ++++++++++++++++++ 5 files changed, 776 insertions(+) create mode 100644 internal/constants/wallet.go create mode 100644 internal/entities/wallet.go create mode 100644 internal/repository/wallet_repository.go create mode 100644 internal/repository/wallet_repository_test.go diff --git a/internal/constants/wallet.go b/internal/constants/wallet.go new file mode 100644 index 0000000..2892fdd --- /dev/null +++ b/internal/constants/wallet.go @@ -0,0 +1,44 @@ +package constants + +// The two balances a customer wallet holds (docs/prd-point-coin.md). EnakPoint pays +// for orders; EnakCoin is spent on games and can be exchanged into EnakPoint. +const ( + WalletCurrencyPoint = "POINT" + WalletCurrencyCoin = "COIN" +) + +func IsValidWalletCurrency(currency string) bool { + return currency == WalletCurrencyPoint || currency == WalletCurrencyCoin +} + +// Ledger row types. §8.1 of the PRD lists, per type, which currency it may use, which +// way it moves the balance, and which reference it must carry. +const ( + WalletTxTypeEarn = "EARN" + WalletTxTypeEarnReversal = "EARN_REVERSAL" + WalletTxTypePayment = "PAYMENT" + WalletTxTypePaymentRefund = "PAYMENT_REFUND" + WalletTxTypeExchangeOut = "EXCHANGE_OUT" + WalletTxTypeExchangeIn = "EXCHANGE_IN" + WalletTxTypeTransferOut = "TRANSFER_OUT" + WalletTxTypeTransferIn = "TRANSFER_IN" + WalletTxTypeGameSpend = "GAME_SPEND" + WalletTxTypeExpire = "EXPIRE" + WalletTxTypeAdjustment = "ADJUSTMENT" + WalletTxTypeMigration = "MIGRATION" + WalletTxTypeRewardRedeem = "REWARD_REDEEM" +) + +// What a ledger row's reference_id points at: where the value came from for a +// credit, or where it went for a debit. +const ( + WalletRefTypeOrder = "ORDER" + WalletRefTypePayment = "PAYMENT" + WalletRefTypeWalletTx = "WALLET_TX" + WalletRefTypeGamePlay = "GAME_PLAY" + WalletRefTypeLot = "LOT" + WalletRefTypeUser = "USER" + WalletRefTypeLegacyPoints = "LEGACY_POINTS" + WalletRefTypeLegacyTokens = "LEGACY_TOKENS" + WalletRefTypeRewardRedemption = "REWARD_REDEMPTION" +) diff --git a/internal/entities/entities.go b/internal/entities/entities.go index a5d6973..e26763b 100644 --- a/internal/entities/entities.go +++ b/internal/entities/entities.go @@ -44,6 +44,11 @@ func GetAllEntities() []interface{} { &ProductOutletPrice{}, &Expense{}, &CashAdvance{}, + // Wallet entities + &CustomerWallet{}, + &WalletTransaction{}, + &WalletLot{}, + &WalletLotAllocation{}, } } diff --git a/internal/entities/wallet.go b/internal/entities/wallet.go new file mode 100644 index 0000000..7b78cb2 --- /dev/null +++ b/internal/entities/wallet.go @@ -0,0 +1,113 @@ +package entities + +import ( + "time" + + "github.com/google/uuid" + "gorm.io/gorm" +) + +// CustomerWallet holds a customer's EnakPoint and EnakCoin balances. The row is also +// the lock every wallet operation for the customer takes first, so concurrent +// operations on one customer queue up instead of spending the same balance twice. +// +// Balances are never written directly: they only move together with a ledger row, and +// only through the wallet processor. +type CustomerWallet struct { + CustomerID uuid.UUID `gorm:"type:uuid;primary_key" json:"customer_id"` + OrganizationID uuid.UUID `gorm:"type:uuid;not null" json:"organization_id"` + PointBalance int64 `gorm:"not null;default:0" json:"point_balance"` + CoinBalance int64 `gorm:"not null;default:0" json:"coin_balance"` + CreatedAt time.Time `gorm:"autoCreateTime" json:"created_at"` + UpdatedAt time.Time `gorm:"autoUpdateTime" json:"updated_at"` +} + +func (CustomerWallet) TableName() string { + return "customer_wallets" +} + +// WalletTransaction is one ledger row. The ledger is append-only: a correction is a +// new row pointing at the one it corrects, never an update. +type WalletTransaction struct { + ID uuid.UUID `gorm:"type:uuid;primary_key;default:gen_random_uuid()" json:"id"` + OrganizationID uuid.UUID `gorm:"type:uuid;not null" json:"organization_id"` + CustomerID uuid.UUID `gorm:"type:uuid;not null" json:"customer_id"` + Currency string `gorm:"not null;size:10" json:"currency"` + Type string `gorm:"not null;size:30" json:"type"` + // Signed: positive credits the wallet, negative debits it. + Amount int64 `gorm:"not null" json:"amount"` + BalanceAfter int64 `gorm:"not null" json:"balance_after"` + GroupID *uuid.UUID `gorm:"type:uuid" json:"group_id"` + + // Where the value came from (credit) or went to (debit). + ReferenceType string `gorm:"not null;size:30" json:"reference_type"` + ReferenceID uuid.UUID `gorm:"type:uuid;not null" json:"reference_id"` + + CounterpartyCustomerID *uuid.UUID `gorm:"type:uuid" json:"counterparty_customer_id"` + ReversesTransactionID *uuid.UUID `gorm:"type:uuid" json:"reverses_transaction_id"` + OutletID *uuid.UUID `gorm:"type:uuid" json:"outlet_id"` + CreatedByUser *uuid.UUID `gorm:"type:uuid" json:"created_by_user"` + Reason *string `gorm:"size:255" json:"reason"` + + // Frozen at creation, so later renames do not rewrite history. + Description string `gorm:"not null;size:255" json:"description"` + Metadata Metadata `gorm:"type:jsonb;default:'{}'" json:"metadata"` + IdempotencyKey *string `gorm:"size:100;unique" json:"idempotency_key"` + CreatedAt time.Time `gorm:"autoCreateTime" json:"created_at"` +} + +func (t *WalletTransaction) BeforeCreate(tx *gorm.DB) error { + if t.ID == uuid.Nil { + t.ID = uuid.New() + } + // A nil map would be stored as JSON null rather than an empty object. + if t.Metadata == nil { + t.Metadata = Metadata{} + } + return nil +} + +func (WalletTransaction) TableName() string { + return "wallet_transactions" +} + +// WalletLot is one credited piece of balance with its own expiry (K9). Debits draw from +// the lots that expire soonest. A lot created by a transfer, exchange or refund carries +// the expiry of the lot it came from and points back at it through OriginLotID. +type WalletLot struct { + ID uuid.UUID `gorm:"type:uuid;primary_key;default:gen_random_uuid()" json:"id"` + OrganizationID uuid.UUID `gorm:"type:uuid;not null" json:"organization_id"` + CustomerID uuid.UUID `gorm:"type:uuid;not null" json:"customer_id"` + Currency string `gorm:"not null;size:10" json:"currency"` + SourceTransactionID uuid.UUID `gorm:"type:uuid;not null" json:"source_transaction_id"` + OriginLotID *uuid.UUID `gorm:"type:uuid" json:"origin_lot_id"` + OriginalAmount int64 `gorm:"not null" json:"original_amount"` + // A cache of OriginalAmount minus the lot's allocations, and the only wallet column + // that is ever updated. + RemainingAmount int64 `gorm:"not null" json:"remaining_amount"` + // Nil means the lot never expires. + ExpiresAt *time.Time `json:"expires_at"` + CreatedAt time.Time `gorm:"autoCreateTime" json:"created_at"` +} + +func (l *WalletLot) BeforeCreate(tx *gorm.DB) error { + if l.ID == uuid.Nil { + l.ID = uuid.New() + } + return nil +} + +func (WalletLot) TableName() string { + return "wallet_lots" +} + +// WalletLotAllocation records how much a debit ledger row drew from one lot. +type WalletLotAllocation struct { + TransactionID uuid.UUID `gorm:"type:uuid;primary_key" json:"transaction_id"` + LotID uuid.UUID `gorm:"type:uuid;primary_key" json:"lot_id"` + Amount int64 `gorm:"not null" json:"amount"` +} + +func (WalletLotAllocation) TableName() string { + return "wallet_lot_allocations" +} diff --git a/internal/repository/wallet_repository.go b/internal/repository/wallet_repository.go new file mode 100644 index 0000000..b33f761 --- /dev/null +++ b/internal/repository/wallet_repository.go @@ -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 +} diff --git a/internal/repository/wallet_repository_test.go b/internal/repository/wallet_repository_test.go new file mode 100644 index 0000000..bfbdd30 --- /dev/null +++ b/internal/repository/wallet_repository_test.go @@ -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) +}