Files
apskel-pos-backend/internal/repository/wallet_repository.go
T
efrilmandClaude Opus 5.5 2eb590caab feat(wallet): add wallet engine as the only way to change a balance
WalletProcessor writes the balance, the ledger row and the lots or
allocations together, which keeps SUM(ledger) = balance = SUM(lot
remaining) (docs/prd-point-coin.md §7.5, PC-104).

- Credit writes the ledger row and creates lots, each with its own expiry
  and origin lot.
- Debit draws from the preferred lots first (a reversal's own lots, or the
  lot being expired), then from unexpired lots in K9 order, and returns the
  allocations with their expiry so CarryOver can give the receiving side of
  a transfer or exchange the same expiry.
- DebitUpTo takes what the wallet has and reports the shortfall (F10, Q3).
- An idempotency key returns the first result; reusing it for a different
  operation is an error.
- §8.1 is checked in code from one rule table, ahead of the database
  constraints, so callers get a readable error.

Each method locks the wallet itself, after validating the input and before
checking the idempotency key, so correctness does not depend on the caller.
Operations on two wallets still call LockWallets first to keep lock order.

Unit tests run on an in-memory repository and check the §7.5 invariants
after every scenario; one more test runs the engine against Postgres when
TEST_DATABASE_URL is set.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
2026-09-30 08:47:37 +07:00

293 lines
11 KiB
Go

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
// GetLotsByIDs returns the lots with the given ids, expired or not, in no
// particular order. Ids that match no lot are left out.
GetLotsByIDs(ctx context.Context, ids []uuid.UUID) ([]entities.WalletLot, error)
// ListLotsBySourceTransaction returns the lots a credit created, oldest first.
ListLotsBySourceTransaction(ctx context.Context, transactionID uuid.UUID) ([]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) GetLotsByIDs(ctx context.Context, ids []uuid.UUID) ([]entities.WalletLot, error) {
var lots []entities.WalletLot
if len(ids) == 0 {
return lots, nil
}
err := DBFromContext(ctx, r.db).WithContext(ctx).
Where("id IN ?", ids).
Find(&lots).Error
if err != nil {
return nil, fmt.Errorf("failed to get wallet lots: %w", err)
}
return lots, nil
}
func (r *walletRepository) ListLotsBySourceTransaction(ctx context.Context, transactionID uuid.UUID) ([]entities.WalletLot, error) {
var lots []entities.WalletLot
err := DBFromContext(ctx, r.db).WithContext(ctx).
Where("source_transaction_id = ?", transactionID).
Order("created_at, id").
Find(&lots).Error
if err != nil {
return nil, fmt.Errorf("failed to list wallet lots by source transaction: %w", err)
}
return lots, nil
}
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
}