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 }