Files
apskel-pos-backend/internal/repository/wallet_repository.go
T
2026-09-30 15:31:11 +07:00

296 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) {
// Find rather than First: a new key is the normal case, and First would log every
// one of them as a "record not found" error.
var walletTxs []entities.WalletTransaction
err := DBFromContext(ctx, r.db).WithContext(ctx).
Where("idempotency_key = ?", key).
Limit(1).
Find(&walletTxs).Error
if err != nil {
return nil, fmt.Errorf("failed to get wallet transaction by idempotency key: %w", err)
}
if len(walletTxs) == 0 {
return nil, nil
}
return &walletTxs[0], 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
}