diff --git a/internal/app/app.go b/internal/app/app.go index ace698f..932aeb4 100644 --- a/internal/app/app.go +++ b/internal/app/app.go @@ -247,6 +247,7 @@ type repositories struct { productOutletPriceRepo *repository.ProductOutletPriceRepositoryImpl expenseRepo *repository.ExpenseRepositoryImpl cashAdvanceRepo *repository.CashAdvanceRepositoryImpl + walletRepo repository.WalletRepository } func (a *App) initRepositories() *repositories { @@ -302,6 +303,7 @@ func (a *App) initRepositories() *repositories { productOutletPriceRepo: repository.NewProductOutletPriceRepositoryImpl(a.db), expenseRepo: repository.NewExpenseRepositoryImpl(a.db), cashAdvanceRepo: repository.NewCashAdvanceRepositoryImpl(a.db), + walletRepo: repository.NewWalletRepository(a.db), } } @@ -350,6 +352,7 @@ type processors struct { productOutletPriceProcessor processor.ProductOutletPriceProcessor expenseProcessor *processor.ExpenseProcessorImpl cashAdvanceProcessor *processor.CashAdvanceProcessorImpl + walletProcessor *processor.WalletProcessor } func (a *App) initProcessors(cfg *config.Config, repos *repositories) *processors { @@ -403,6 +406,7 @@ func (a *App) initProcessors(cfg *config.Config, repos *repositories) *processor productOutletPriceProcessor: processor.NewProductOutletPriceProcessorImpl(repos.productOutletPriceRepo, repos.productRepo, repos.outletRepo), expenseProcessor: processor.NewExpenseProcessorImpl(repos.expenseRepo, repos.purchaseCategoryRepo, repos.cashAdvanceRepo), cashAdvanceProcessor: processor.NewCashAdvanceProcessorImpl(repos.cashAdvanceRepo, repos.categoryRepo), + walletProcessor: processor.NewWalletProcessor(repos.walletRepo), } } diff --git a/internal/processor/wallet_processor.go b/internal/processor/wallet_processor.go new file mode 100644 index 0000000..3190fa3 --- /dev/null +++ b/internal/processor/wallet_processor.go @@ -0,0 +1,527 @@ +package processor + +import ( + "context" + "errors" + "fmt" + "strings" + "time" + + "github.com/google/uuid" + + "apskel-pos-be/internal/constants" + "apskel-pos-be/internal/entities" + "apskel-pos-be/internal/repository" +) + +var ( + // ErrWalletInvalidEntry wraps every rejection of an entry that breaks the rules in + // docs/prd-point-coin.md §8.1. The database enforces most of them too; checking + // here first gives callers a readable error instead of a constraint name. + ErrWalletInvalidEntry = errors.New("wallet: invalid entry") + // ErrWalletIdempotencyConflict means an idempotency key was reused for a different + // operation. Retrying the same operation with the same key is not a conflict. + ErrWalletIdempotencyConflict = errors.New("wallet: idempotency key already used for a different operation") +) + +// WalletEntry is what every ledger row needs, whichever way it moves the balance. +// Which of the optional fields a type requires is listed in §8.1. +type WalletEntry struct { + // Optional. Set it when another row must reference this one before it exists, as + // the two rows of an exchange or a transfer do. + TransactionID uuid.UUID + + CustomerID uuid.UUID + Currency string + Type string + // Always positive: Credit adds it, Debit takes it away. + Amount int64 + + ReferenceType string + ReferenceID uuid.UUID + + GroupID *uuid.UUID + CounterpartyCustomerID *uuid.UUID + ReversesTransactionID *uuid.UUID + OutletID *uuid.UUID + CreatedByUser *uuid.UUID + Reason *string + + Description string + Metadata entities.Metadata + // Optional. A retry with the same key returns the first result without moving + // anything again. + IdempotencyKey string +} + +// WalletLotInput is one lot a credit creates. +type WalletLotInput struct { + Amount int64 + // Nil means the lot never expires. + ExpiresAt *time.Time + // The lot this one was carried over from, for transfers, exchanges and refunds. + OriginLotID *uuid.UUID +} + +type WalletCreditInput struct { + WalletEntry + // How the credit is split into lots. Their amounts must add up to Amount. Leave + // empty for a single lot that never expires. + Lots []WalletLotInput +} + +type WalletDebitInput struct { + WalletEntry + // Lots to draw from first, in this order, before falling back to the K9 order. + // A reversal names the lots its EARN created (F10), and the expiry job names the + // lot that expired. These lots are used even if they have already expired. + PreferredLotIDs []uuid.UUID +} + +// WalletAllocation is how much a debit took from one lot. It carries the lot's +// expiry, so a transfer or exchange can give the receiving lot the same expiry (K9). +type WalletAllocation struct { + LotID uuid.UUID + Amount int64 + ExpiresAt *time.Time +} + +type WalletResult struct { + // Nil only when DebitUpTo found nothing to take. + Transaction *entities.WalletTransaction + // The lots a credit created. + Lots []entities.WalletLot + // The lots a debit drew from, in the order they were used. + Allocations []WalletAllocation + // What DebitUpTo could not take because the balance ran out. + Shortfall int64 + // True when the idempotency key had already been used and nothing moved. + Replayed bool +} + +// CarryOver turns a debit's allocations into lots for the receiving side of a +// transfer or exchange. Each lot keeps the expiry of the lot it came from and points +// back at it, so a balance cannot be kept alive by moving it around (K9). +func (r *WalletResult) CarryOver() []WalletLotInput { + lots := make([]WalletLotInput, 0, len(r.Allocations)) + for _, a := range r.Allocations { + lotID := a.LotID + lots = append(lots, WalletLotInput{Amount: a.Amount, ExpiresAt: a.ExpiresAt, OriginLotID: &lotID}) + } + return lots +} + +// WalletProcessor is the only code allowed to change a wallet balance. Every change +// writes the balance, the ledger row and the lots or allocations together, which is +// what keeps SUM(ledger) = balance = SUM(lot remaining) (§7.5). +// +// Every method must run inside a transaction from TxManager, and the repository +// refuses otherwise. Each method locks the customer's wallet itself, so a single-wallet +// caller needs nothing more. A caller touching two wallets, such as a transfer, must +// call LockWallets first so the locks are always taken in the same order. +type WalletProcessor struct { + repo repository.WalletRepository + now func() time.Time +} + +func NewWalletProcessor(repo repository.WalletRepository) *WalletProcessor { + return &WalletProcessor{repo: repo, now: time.Now} +} + +// LockWallet locks one customer's wallet, creating it if needed. Credit and Debit do +// this themselves; call it when something must be read under the lock first. +func (p *WalletProcessor) LockWallet(ctx context.Context, customerID uuid.UUID) error { + _, err := p.repo.LockWallet(ctx, customerID) + return err +} + +// LockWallets locks two customers' wallets in a fixed order. Call it before touching +// both wallets in one transaction. +func (p *WalletProcessor) LockWallets(ctx context.Context, a, b uuid.UUID) error { + _, _, err := p.repo.LockWallets(ctx, a, b) + return err +} + +// Credit adds Amount to the wallet and creates its lots. +func (p *WalletProcessor) Credit(ctx context.Context, in WalletCreditInput) (*WalletResult, error) { + if err := validateWalletEntry(&in.WalletEntry, true); err != nil { + return nil, err + } + lots := in.Lots + if len(lots) == 0 { + lots = []WalletLotInput{{Amount: in.Amount}} + } + var total int64 + for _, lot := range lots { + if lot.Amount <= 0 { + return nil, fmt.Errorf("%w: lot amount must be positive, got %d", ErrWalletInvalidEntry, lot.Amount) + } + total += lot.Amount + } + if total != in.Amount { + return nil, fmt.Errorf("%w: lots add up to %d, not %d", ErrWalletInvalidEntry, total, in.Amount) + } + + wallet, err := p.repo.LockWallet(ctx, in.CustomerID) + if err != nil { + return nil, err + } + if replay, err := p.replay(ctx, &in.WalletEntry, true, true); replay != nil || err != nil { + return replay, err + } + + balance, err := p.repo.AddBalance(ctx, in.CustomerID, in.Currency, in.Amount) + if err != nil { + return nil, err + } + walletTx := newWalletTransaction(wallet, &in.WalletEntry, in.Amount, balance, nil) + if err := p.repo.CreateTransaction(ctx, walletTx); err != nil { + return nil, fmt.Errorf("failed to create wallet transaction: %w", err) + } + + result := &WalletResult{Transaction: walletTx} + for _, lotIn := range lots { + lot := entities.WalletLot{ + OrganizationID: wallet.OrganizationID, + CustomerID: in.CustomerID, + Currency: in.Currency, + SourceTransactionID: walletTx.ID, + OriginLotID: lotIn.OriginLotID, + OriginalAmount: lotIn.Amount, + RemainingAmount: lotIn.Amount, + ExpiresAt: lotIn.ExpiresAt, + } + if err := p.repo.CreateLot(ctx, &lot); err != nil { + return nil, fmt.Errorf("failed to create wallet lot: %w", err) + } + result.Lots = append(result.Lots, lot) + } + return result, nil +} + +// Debit takes exactly Amount from the wallet, or nothing at all with +// repository.ErrWalletInsufficientBalance if the usable balance is short. +func (p *WalletProcessor) Debit(ctx context.Context, in WalletDebitInput) (*WalletResult, error) { + return p.debit(ctx, in, false) +} + +// DebitUpTo takes as much of Amount as the wallet has and reports the rest as +// Shortfall. It is for reversing earnings the customer has already spent (F10, Q3). +// When there is nothing to take, no ledger row is written and Transaction is nil; +// such a call leaves no trace, so a retry with the same key takes whatever the +// balance holds by then. +func (p *WalletProcessor) DebitUpTo(ctx context.Context, in WalletDebitInput) (*WalletResult, error) { + return p.debit(ctx, in, true) +} + +func (p *WalletProcessor) debit(ctx context.Context, in WalletDebitInput, upTo bool) (*WalletResult, error) { + if err := validateWalletEntry(&in.WalletEntry, false); err != nil { + return nil, err + } + wallet, err := p.repo.LockWallet(ctx, in.CustomerID) + if err != nil { + return nil, err + } + // DebitUpTo may have taken less than asked, so the amount cannot be compared. + if replay, err := p.replay(ctx, &in.WalletEntry, false, !upTo); replay != nil || err != nil { + return replay, err + } + + lots, err := p.spendableLots(ctx, &in) + if err != nil { + return nil, err + } + + var available int64 + for _, lot := range lots { + available += lot.RemainingAmount + } + take := in.Amount + if available < take { + if !upTo { + return nil, repository.ErrWalletInsufficientBalance + } + take = available + } + result := &WalletResult{Shortfall: in.Amount - take} + if take == 0 { + return result, nil + } + + var metadata entities.Metadata + if upTo { + metadata = entities.Metadata{"requested_amount": in.Amount, "shortfall": result.Shortfall} + } + + balance, err := p.repo.AddBalance(ctx, in.CustomerID, in.Currency, -take) + if err != nil { + return nil, err + } + walletTx := newWalletTransaction(wallet, &in.WalletEntry, -take, balance, metadata) + if err := p.repo.CreateTransaction(ctx, walletTx); err != nil { + return nil, fmt.Errorf("failed to create wallet transaction: %w", err) + } + result.Transaction = walletTx + + var allocations []entities.WalletLotAllocation + remaining := take + for _, lot := range lots { + if remaining == 0 { + break + } + amount := min(lot.RemainingAmount, remaining) + remaining -= amount + if err := p.repo.ConsumeLot(ctx, lot.ID, amount); err != nil { + return nil, err + } + allocations = append(allocations, entities.WalletLotAllocation{TransactionID: walletTx.ID, LotID: lot.ID, Amount: amount}) + result.Allocations = append(result.Allocations, WalletAllocation{LotID: lot.ID, Amount: amount, ExpiresAt: lot.ExpiresAt}) + } + if err := p.repo.CreateAllocations(ctx, allocations); err != nil { + return nil, fmt.Errorf("failed to create wallet lot allocations: %w", err) + } + return result, nil +} + +// spendableLots returns the lots a debit may draw from, in the order it draws: the +// preferred lots first, then the unexpired lots in K9 order. +func (p *WalletProcessor) spendableLots(ctx context.Context, in *WalletDebitInput) ([]entities.WalletLot, error) { + var lots []entities.WalletLot + preferred := make(map[uuid.UUID]bool, len(in.PreferredLotIDs)) + + if len(in.PreferredLotIDs) > 0 { + found, err := p.repo.GetLotsByIDs(ctx, in.PreferredLotIDs) + if err != nil { + return nil, err + } + byID := make(map[uuid.UUID]entities.WalletLot, len(found)) + for _, lot := range found { + byID[lot.ID] = lot + } + for _, id := range in.PreferredLotIDs { + lot, ok := byID[id] + if !ok || lot.CustomerID != in.CustomerID || lot.Currency != in.Currency { + return nil, fmt.Errorf("%w: lot %s is not a %s lot of this customer", ErrWalletInvalidEntry, id, in.Currency) + } + if preferred[id] { + continue + } + preferred[id] = true + if lot.RemainingAmount > 0 { + lots = append(lots, lot) + } + } + } + + active, err := p.repo.ListActiveLots(ctx, in.CustomerID, in.Currency, p.now()) + if err != nil { + return nil, err + } + for _, lot := range active { + if !preferred[lot.ID] { + lots = append(lots, lot) + } + } + return lots, nil +} + +// replay returns the first result for an idempotency key that has already been used, +// or nil when the key is new. It runs after the wallet lock, so a concurrent request +// with the same key has either committed its row or not started. +func (p *WalletProcessor) replay(ctx context.Context, in *WalletEntry, credit, compareAmount bool) (*WalletResult, error) { + if in.IdempotencyKey == "" { + return nil, nil + } + walletTx, err := p.repo.GetTransactionByIdempotencyKey(ctx, in.IdempotencyKey) + if err != nil || walletTx == nil { + return nil, err + } + + sameDirection := (walletTx.Amount > 0) == credit + sameAmount := !compareAmount || abs(walletTx.Amount) == in.Amount + if walletTx.CustomerID != in.CustomerID || walletTx.Currency != in.Currency || + walletTx.Type != in.Type || !sameDirection || !sameAmount { + return nil, ErrWalletIdempotencyConflict + } + + result := &WalletResult{Transaction: walletTx, Replayed: true} + if credit { + result.Lots, err = p.repo.ListLotsBySourceTransaction(ctx, walletTx.ID) + return result, err + } + + // JSON numbers come back from JSONB as float64. + switch shortfall := walletTx.Metadata["shortfall"].(type) { + case float64: + result.Shortfall = int64(shortfall) + case int64: + result.Shortfall = shortfall + } + allocations, err := p.repo.ListAllocationsByTransaction(ctx, walletTx.ID) + if err != nil { + return nil, err + } + ids := make([]uuid.UUID, 0, len(allocations)) + for _, a := range allocations { + ids = append(ids, a.LotID) + } + lots, err := p.repo.GetLotsByIDs(ctx, ids) + if err != nil { + return nil, err + } + expiry := make(map[uuid.UUID]*time.Time, len(lots)) + for _, lot := range lots { + expiry[lot.ID] = lot.ExpiresAt + } + for _, a := range allocations { + result.Allocations = append(result.Allocations, WalletAllocation{LotID: a.LotID, Amount: a.Amount, ExpiresAt: expiry[a.LotID]}) + } + return result, nil +} + +func newWalletTransaction(wallet *entities.CustomerWallet, in *WalletEntry, amount, balance int64, extra entities.Metadata) *entities.WalletTransaction { + metadata := entities.Metadata{} + for k, v := range in.Metadata { + metadata[k] = v + } + for k, v := range extra { + metadata[k] = v + } + var key *string + if in.IdempotencyKey != "" { + k := in.IdempotencyKey + key = &k + } + return &entities.WalletTransaction{ + ID: in.TransactionID, + OrganizationID: wallet.OrganizationID, + CustomerID: in.CustomerID, + Currency: in.Currency, + Type: in.Type, + Amount: amount, + BalanceAfter: balance, + GroupID: in.GroupID, + ReferenceType: in.ReferenceType, + ReferenceID: in.ReferenceID, + CounterpartyCustomerID: in.CounterpartyCustomerID, + ReversesTransactionID: in.ReversesTransactionID, + OutletID: in.OutletID, + CreatedByUser: in.CreatedByUser, + Reason: in.Reason, + Description: in.Description, + Metadata: metadata, + IdempotencyKey: key, + } +} + +// walletTypeRule is one row of §8.1. +type walletTypeRule struct { + credit, debit bool + currency string // empty: either currency + referenceTypes []string + needsOutlet bool + needsReverses bool + needsGroup bool + needsCounter bool + needsActor bool +} + +var walletTypeRules = map[string]walletTypeRule{ + constants.WalletTxTypeEarn: {credit: true, referenceTypes: []string{constants.WalletRefTypeOrder}, needsOutlet: true}, + constants.WalletTxTypeEarnReversal: {debit: true, referenceTypes: []string{constants.WalletRefTypeOrder}, needsOutlet: true, needsReverses: true}, + constants.WalletTxTypePayment: {debit: true, currency: constants.WalletCurrencyPoint, referenceTypes: []string{constants.WalletRefTypePayment}, needsOutlet: true}, + constants.WalletTxTypePaymentRefund: {credit: true, currency: constants.WalletCurrencyPoint, referenceTypes: []string{constants.WalletRefTypePayment}, needsOutlet: true, needsReverses: true}, + constants.WalletTxTypeExchangeOut: {debit: true, currency: constants.WalletCurrencyCoin, referenceTypes: []string{constants.WalletRefTypeWalletTx}, needsGroup: true}, + constants.WalletTxTypeExchangeIn: {credit: true, currency: constants.WalletCurrencyPoint, referenceTypes: []string{constants.WalletRefTypeWalletTx}, needsGroup: true}, + constants.WalletTxTypeTransferOut: {debit: true, referenceTypes: []string{constants.WalletRefTypeWalletTx}, needsGroup: true, needsCounter: true}, + constants.WalletTxTypeTransferIn: {credit: true, referenceTypes: []string{constants.WalletRefTypeWalletTx}, needsGroup: true, needsCounter: true}, + constants.WalletTxTypeGameSpend: {debit: true, currency: constants.WalletCurrencyCoin, referenceTypes: []string{constants.WalletRefTypeGamePlay}}, + constants.WalletTxTypeExpire: {debit: true, referenceTypes: []string{constants.WalletRefTypeLot}}, + constants.WalletTxTypeAdjustment: {credit: true, debit: true, referenceTypes: []string{constants.WalletRefTypeUser}, needsActor: true}, + constants.WalletTxTypeMigration: {credit: true, referenceTypes: []string{constants.WalletRefTypeLegacyPoints, constants.WalletRefTypeLegacyTokens}}, + constants.WalletTxTypeRewardRedeem: {debit: true, currency: constants.WalletCurrencyPoint, referenceTypes: []string{constants.WalletRefTypeRewardRedemption}}, +} + +func validateWalletEntry(in *WalletEntry, credit bool) error { + invalid := func(format string, args ...any) error { + return fmt.Errorf("%w: %s", ErrWalletInvalidEntry, fmt.Sprintf(format, args...)) + } + + rule, ok := walletTypeRules[in.Type] + if !ok { + return invalid("unknown type %q", in.Type) + } + if credit && !rule.credit { + return invalid("%s cannot add to a balance", in.Type) + } + if !credit && !rule.debit { + return invalid("%s cannot take from a balance", in.Type) + } + if in.CustomerID == uuid.Nil { + return invalid("customer is required") + } + if !constants.IsValidWalletCurrency(in.Currency) { + return invalid("unknown currency %q", in.Currency) + } + if rule.currency != "" && in.Currency != rule.currency { + return invalid("%s must be in %s", in.Type, rule.currency) + } + if in.Amount <= 0 { + return invalid("amount must be positive, got %d", in.Amount) + } + if !containsString(rule.referenceTypes, in.ReferenceType) { + return invalid("%s must reference %s, got %q", in.Type, strings.Join(rule.referenceTypes, " or "), in.ReferenceType) + } + if in.ReferenceID == uuid.Nil { + return invalid("reference id is required") + } + if strings.TrimSpace(in.Description) == "" { + return invalid("description is required") + } + if rule.needsOutlet && isNilID(in.OutletID) { + return invalid("%s requires an outlet", in.Type) + } + if rule.needsReverses && isNilID(in.ReversesTransactionID) { + return invalid("%s requires the transaction it reverses", in.Type) + } + if rule.needsGroup && isNilID(in.GroupID) { + return invalid("%s requires a group id", in.Type) + } + if rule.needsCounter { + if isNilID(in.CounterpartyCustomerID) { + return invalid("%s requires a counterparty", in.Type) + } + if *in.CounterpartyCustomerID == in.CustomerID { + return invalid("%s cannot go to the same customer", in.Type) + } + } + if rule.needsActor { + if isNilID(in.CreatedByUser) { + return invalid("%s requires the admin who made it", in.Type) + } + if in.Reason == nil || strings.TrimSpace(*in.Reason) == "" { + return invalid("%s requires a reason", in.Type) + } + } + return nil +} + +func isNilID(id *uuid.UUID) bool { + return id == nil || *id == uuid.Nil +} + +func containsString(values []string, v string) bool { + for _, value := range values { + if value == v { + return true + } + } + return false +} + +func abs(v int64) int64 { + if v < 0 { + return -v + } + return v +} diff --git a/internal/processor/wallet_processor_db_test.go b/internal/processor/wallet_processor_db_test.go new file mode 100644 index 0000000..d3aa795 --- /dev/null +++ b/internal/processor/wallet_processor_db_test.go @@ -0,0 +1,142 @@ +package processor + +import ( + "context" + "os" + "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/repository" +) + +// Runs the engine against Postgres, to show the rows it writes pass the database +// constraints and reconcile the way §7.5 requires. Needs TEST_DATABASE_URL pointing +// at a migrated database; see internal/repository/wallet_repository_test.go. +func TestWalletProcessor_AgainstPostgres(t *testing.T) { + 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) + + org, a, b := uuid.New(), uuid.New(), uuid.New() + require.NoError(t, db.Exec(`INSERT INTO organizations (id, name, plan_type) VALUES (?, 'wallet test', 'basic')`, org).Error) + require.NoError(t, db.Exec(`INSERT INTO customers (id, organization_id, name) VALUES (?, ?, 'A'), (?, ?, 'B')`, a, org, b, org).Error) + customers := []uuid.UUID{a, b} + t.Cleanup(func() { + db.Exec(`DELETE FROM wallet_lot_allocations WHERE lot_id IN (SELECT id FROM wallet_lots WHERE customer_id IN ?)`, customers) + db.Exec(`DELETE FROM wallet_lots WHERE customer_id IN ?`, customers) + db.Exec(`DELETE FROM wallet_transactions WHERE customer_id IN ?`, customers) + db.Exec(`DELETE FROM customer_wallets WHERE customer_id IN ?`, customers) + db.Exec(`DELETE FROM customers WHERE id IN ?`, customers) + db.Exec(`DELETE FROM organizations WHERE id = ?`, org) + }) + + p := NewWalletProcessor(repository.NewWalletRepository(db)) + txm := repository.NewTxManager(db) + now := time.Now() + inTx := func(fn func(ctx context.Context) error) { + t.Helper() + require.NoError(t, txm.WithTransaction(context.Background(), fn)) + } + + // Without a transaction nothing moves. + _, err = p.Credit(context.Background(), earn(a, 10, nil)) + assert.ErrorIs(t, err, repository.ErrWalletTxRequired) + + var earned *WalletResult + inTx(func(ctx context.Context) error { + soon := now.Add(time.Hour) + earned, err = p.Credit(ctx, earn(a, 100, &soon)) + require.NoError(t, err) + _, err = p.Credit(ctx, earn(a, 50, nil)) + return err + }) + + // Transfer 120 from A to B, spanning both of A's lots. + inTx(func(ctx context.Context) error { + require.NoError(t, p.LockWallets(ctx, a, b)) + group, outID, inID := uuid.New(), uuid.New(), uuid.New() + out, err := p.Debit(ctx, WalletDebitInput{WalletEntry: WalletEntry{ + TransactionID: outID, CustomerID: a, Currency: constants.WalletCurrencyPoint, + Type: constants.WalletTxTypeTransferOut, Amount: 120, + ReferenceType: constants.WalletRefTypeWalletTx, ReferenceID: inID, + GroupID: &group, CounterpartyCustomerID: &b, Description: "Transfer ke B", + }}) + require.NoError(t, err) + _, err = p.Credit(ctx, WalletCreditInput{ + WalletEntry: WalletEntry{ + TransactionID: inID, CustomerID: b, Currency: constants.WalletCurrencyPoint, + Type: constants.WalletTxTypeTransferIn, Amount: 120, + ReferenceType: constants.WalletRefTypeWalletTx, ReferenceID: outID, + GroupID: &group, CounterpartyCustomerID: &a, Description: "Transfer dari A", + }, + Lots: out.CarryOver(), + }) + return err + }) + + // Reversing the 100 earned leaves A 30 short. The retry reads the shortfall back + // out of JSONB and takes nothing more. + rev := reversal(a, 100, earned) + rev.IdempotencyKey = "reverse:" + earned.Transaction.ID.String() + var first, second *WalletResult + inTx(func(ctx context.Context) error { + first, err = p.DebitUpTo(ctx, rev) + return err + }) + inTx(func(ctx context.Context) error { + second, err = p.DebitUpTo(ctx, rev) + return err + }) + assert.Equal(t, int64(-30), first.Transaction.Amount) + assert.Equal(t, int64(70), first.Shortfall) + assert.True(t, second.Replayed) + assert.Equal(t, int64(70), second.Shortfall) + assert.Equal(t, first.Transaction.ID, second.Transaction.ID) + + // Overdraw fails and rolls back cleanly. + err = txm.WithTransaction(context.Background(), func(ctx context.Context) error { + _, err := p.Debit(ctx, pay(b, 121)) + return err + }) + assert.ErrorIs(t, err, repository.ErrWalletInsufficientBalance) + + var balances []struct { + CustomerID uuid.UUID + PointBalance int64 + } + require.NoError(t, db.Raw(`SELECT customer_id, point_balance FROM customer_wallets WHERE customer_id IN ?`, customers).Scan(&balances).Error) + got := map[uuid.UUID]int64{} + for _, row := range balances { + got[row.CustomerID] = row.PointBalance + } + assert.Equal(t, map[uuid.UUID]int64{a: 0, b: 120}, got) + + // §7.5, straight from the tables. + var broken []string + require.NoError(t, db.Raw(` + SELECT 'wallet ' || w.customer_id FROM customer_wallets w + WHERE w.customer_id IN ? AND ( + w.point_balance <> (SELECT COALESCE(SUM(amount), 0) FROM wallet_transactions t WHERE t.customer_id = w.customer_id AND t.currency = 'POINT') + OR w.point_balance <> (SELECT COALESCE(SUM(remaining_amount), 0) FROM wallet_lots l WHERE l.customer_id = w.customer_id AND l.currency = 'POINT')) + UNION ALL + SELECT 'lot ' || l.id FROM wallet_lots l + WHERE l.customer_id IN ? AND l.original_amount - l.remaining_amount + <> (SELECT COALESCE(SUM(amount), 0) FROM wallet_lot_allocations a WHERE a.lot_id = l.id) + UNION ALL + SELECT 'debit ' || t.id FROM wallet_transactions t + WHERE t.customer_id IN ? AND t.amount < 0 + AND -t.amount <> (SELECT COALESCE(SUM(amount), 0) FROM wallet_lot_allocations a WHERE a.transaction_id = t.id)`, + customers, customers, customers).Scan(&broken).Error) + assert.Empty(t, broken) +} diff --git a/internal/processor/wallet_processor_test.go b/internal/processor/wallet_processor_test.go new file mode 100644 index 0000000..810dfd7 --- /dev/null +++ b/internal/processor/wallet_processor_test.go @@ -0,0 +1,788 @@ +package processor + +import ( + "context" + "errors" + "sort" + "testing" + "time" + + "github.com/google/uuid" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "apskel-pos-be/internal/constants" + "apskel-pos-be/internal/entities" + "apskel-pos-be/internal/repository" +) + +// walletRepoFake is an in-memory WalletRepository with the same conditional-update +// semantics as the real one. Locks are only counted: these tests are single-threaded, +// and the locking itself is covered by the repository tests against Postgres. +type walletRepoFake struct { + customers map[uuid.UUID]uuid.UUID // customer -> organization + wallets map[uuid.UUID]*entities.CustomerWallet + transactions []*entities.WalletTransaction + lots []*entities.WalletLot + allocations []entities.WalletLotAllocation + locks map[uuid.UUID]int + clock time.Time +} + +func newWalletRepoFake() *walletRepoFake { + return &walletRepoFake{ + customers: map[uuid.UUID]uuid.UUID{}, + wallets: map[uuid.UUID]*entities.CustomerWallet{}, + locks: map[uuid.UUID]int{}, + clock: time.Date(2026, 1, 1, 0, 0, 0, 0, time.UTC), + } +} + +func (f *walletRepoFake) tick() time.Time { + f.clock = f.clock.Add(time.Second) + return f.clock +} + +func (f *walletRepoFake) LockWallet(_ context.Context, customerID uuid.UUID) (*entities.CustomerWallet, error) { + org, ok := f.customers[customerID] + if !ok { + return nil, repository.ErrWalletNotFound + } + if f.wallets[customerID] == nil { + f.wallets[customerID] = &entities.CustomerWallet{CustomerID: customerID, OrganizationID: org} + } + f.locks[customerID]++ + w := *f.wallets[customerID] + return &w, nil +} + +func (f *walletRepoFake) LockWallets(ctx context.Context, a, b uuid.UUID) (*entities.CustomerWallet, *entities.CustomerWallet, error) { + wa, err := f.LockWallet(ctx, a) + if err != nil { + return nil, nil, err + } + wb, err := f.LockWallet(ctx, b) + return wa, wb, err +} + +func (f *walletRepoFake) AddBalance(_ context.Context, customerID uuid.UUID, currency string, delta int64) (int64, error) { + w := f.wallets[customerID] + if w == nil { + return 0, repository.ErrWalletNotFound + } + balance := &w.PointBalance + if currency == constants.WalletCurrencyCoin { + balance = &w.CoinBalance + } + if *balance+delta < 0 { + return 0, repository.ErrWalletInsufficientBalance + } + *balance += delta + return *balance, nil +} + +func (f *walletRepoFake) GetWallet(_ context.Context, customerID uuid.UUID) (*entities.CustomerWallet, error) { + w := f.wallets[customerID] + if w == nil { + return nil, errors.New("not found") + } + c := *w + return &c, nil +} + +func (f *walletRepoFake) CreateTransaction(_ context.Context, tx *entities.WalletTransaction) error { + if tx.IdempotencyKey != nil { + for _, t := range f.transactions { + if t.IdempotencyKey != nil && *t.IdempotencyKey == *tx.IdempotencyKey { + return errors.New("duplicate idempotency key") + } + } + } + if tx.ID == uuid.Nil { + tx.ID = uuid.New() + } + tx.CreatedAt = f.tick() + c := *tx + f.transactions = append(f.transactions, &c) + return nil +} + +func (f *walletRepoFake) GetTransactionByIdempotencyKey(_ context.Context, key string) (*entities.WalletTransaction, error) { + for _, t := range f.transactions { + if t.IdempotencyKey != nil && *t.IdempotencyKey == key { + c := *t + return &c, nil + } + } + return nil, nil +} + +func (f *walletRepoFake) CreateLot(_ context.Context, lot *entities.WalletLot) error { + if lot.ID == uuid.Nil { + lot.ID = uuid.New() + } + lot.CreatedAt = f.tick() + c := *lot + f.lots = append(f.lots, &c) + return nil +} + +func (f *walletRepoFake) GetLotsByIDs(_ context.Context, ids []uuid.UUID) ([]entities.WalletLot, error) { + var out []entities.WalletLot + for _, lot := range f.lots { + for _, id := range ids { + if lot.ID == id { + out = append(out, *lot) + break + } + } + } + return out, nil +} + +func (f *walletRepoFake) ListLotsBySourceTransaction(_ context.Context, txID uuid.UUID) ([]entities.WalletLot, error) { + var out []entities.WalletLot + for _, lot := range f.lots { + if lot.SourceTransactionID == txID { + out = append(out, *lot) + } + } + return out, nil +} + +func (f *walletRepoFake) ListActiveLots(_ context.Context, customerID uuid.UUID, currency string, asOf time.Time) ([]entities.WalletLot, error) { + var out []entities.WalletLot + for _, lot := range f.lots { + if lot.CustomerID == customerID && lot.Currency == currency && lot.RemainingAmount > 0 && + (lot.ExpiresAt == nil || lot.ExpiresAt.After(asOf)) { + out = append(out, *lot) + } + } + sort.SliceStable(out, func(i, j int) bool { + a, b := out[i], out[j] + switch { + case a.ExpiresAt == nil && b.ExpiresAt != nil: + return false + case a.ExpiresAt != nil && b.ExpiresAt == nil: + return true + case a.ExpiresAt != nil && !a.ExpiresAt.Equal(*b.ExpiresAt): + return a.ExpiresAt.Before(*b.ExpiresAt) + } + return a.CreatedAt.Before(b.CreatedAt) + }) + return out, nil +} + +func (f *walletRepoFake) ConsumeLot(_ context.Context, lotID uuid.UUID, amount int64) error { + for _, lot := range f.lots { + if lot.ID == lotID { + if lot.RemainingAmount < amount { + return repository.ErrWalletLotInsufficient + } + lot.RemainingAmount -= amount + return nil + } + } + return repository.ErrWalletLotInsufficient +} + +func (f *walletRepoFake) CreateAllocations(_ context.Context, allocations []entities.WalletLotAllocation) error { + f.allocations = append(f.allocations, allocations...) + return nil +} + +func (f *walletRepoFake) ListAllocationsByTransaction(_ context.Context, txID uuid.UUID) ([]entities.WalletLotAllocation, error) { + var out []entities.WalletLotAllocation + for _, a := range f.allocations { + if a.TransactionID == txID { + out = append(out, a) + } + } + return out, nil +} + +// assertInvariants checks the reconciliation rules of §7.5 over everything the fake +// holds. +func (f *walletRepoFake) assertInvariants(t *testing.T) { + t.Helper() + allocatedFromLot := map[uuid.UUID]int64{} + allocatedByTx := map[uuid.UUID]int64{} + for _, a := range f.allocations { + allocatedFromLot[a.LotID] += a.Amount + allocatedByTx[a.TransactionID] += a.Amount + } + createdByTx := map[uuid.UUID]int64{} + for _, lot := range f.lots { + createdByTx[lot.SourceTransactionID] += lot.OriginalAmount + assert.Equal(t, lot.OriginalAmount-allocatedFromLot[lot.ID], lot.RemainingAmount, "lot %s: original - allocations = remaining", lot.ID) + } + for _, tx := range f.transactions { + if tx.Amount > 0 { + assert.Equal(t, tx.Amount, createdByTx[tx.ID], "credit %s: lots add up to the amount", tx.Type) + assert.Zero(t, allocatedByTx[tx.ID], "credit %s has no allocations", tx.Type) + } else { + assert.Equal(t, -tx.Amount, allocatedByTx[tx.ID], "debit %s: allocations add up to the amount", tx.Type) + assert.Zero(t, createdByTx[tx.ID], "debit %s creates no lots", tx.Type) + } + } + for customerID, w := range f.wallets { + for currency, balance := range map[string]int64{ + constants.WalletCurrencyPoint: w.PointBalance, + constants.WalletCurrencyCoin: w.CoinBalance, + } { + var ledger, lots, last int64 + for _, tx := range f.transactions { + if tx.CustomerID == customerID && tx.Currency == currency { + ledger += tx.Amount + last = tx.BalanceAfter + } + } + for _, lot := range f.lots { + if lot.CustomerID == customerID && lot.Currency == currency { + lots += lot.RemainingAmount + } + } + assert.Equal(t, balance, ledger, "%s balance = SUM(ledger)", currency) + assert.Equal(t, balance, lots, "%s balance = SUM(lot remaining)", currency) + assert.Equal(t, balance, last, "%s balance = last balance_after", currency) + } + } +} + +type walletTestEnv struct { + repo *walletRepoFake + p *WalletProcessor + now time.Time + org uuid.UUID + ctx context.Context +} + +func newWalletTestEnv(t *testing.T) *walletTestEnv { + repo := newWalletRepoFake() + env := &walletTestEnv{ + repo: repo, + p: NewWalletProcessor(repo), + now: time.Date(2026, 6, 1, 12, 0, 0, 0, time.UTC), + org: uuid.New(), + ctx: context.Background(), + } + env.p.now = func() time.Time { return env.now } + t.Cleanup(func() { repo.assertInvariants(t) }) + return env +} + +func (e *walletTestEnv) customer() uuid.UUID { + id := uuid.New() + e.repo.customers[id] = e.org + return id +} + +func (e *walletTestEnv) at(d time.Duration) *time.Time { + v := e.now.Add(d) + return &v +} + +func ptr[T any](v T) *T { return &v } + +func earn(customerID uuid.UUID, amount int64, expiresAt *time.Time) WalletCreditInput { + return WalletCreditInput{ + WalletEntry: WalletEntry{ + CustomerID: customerID, + Currency: constants.WalletCurrencyPoint, + Type: constants.WalletTxTypeEarn, + Amount: amount, + ReferenceType: constants.WalletRefTypeOrder, + ReferenceID: uuid.New(), + OutletID: ptr(uuid.New()), + Description: "Belanja", + }, + Lots: []WalletLotInput{{Amount: amount, ExpiresAt: expiresAt}}, + } +} + +func pay(customerID uuid.UUID, amount int64) WalletDebitInput { + return WalletDebitInput{WalletEntry: WalletEntry{ + CustomerID: customerID, + Currency: constants.WalletCurrencyPoint, + Type: constants.WalletTxTypePayment, + Amount: amount, + ReferenceType: constants.WalletRefTypePayment, + ReferenceID: uuid.New(), + OutletID: ptr(uuid.New()), + Description: "Bayar", + }} +} + +func (e *walletTestEnv) credit(t *testing.T, in WalletCreditInput) *WalletResult { + t.Helper() + res, err := e.p.Credit(e.ctx, in) + require.NoError(t, err) + return res +} + +func (e *walletTestEnv) balance(t *testing.T, customerID uuid.UUID) int64 { + t.Helper() + w, err := e.repo.GetWallet(e.ctx, customerID) + require.NoError(t, err) + return w.PointBalance +} + +func allocationsOf(res *WalletResult) map[uuid.UUID]int64 { + out := map[uuid.UUID]int64{} + for _, a := range res.Allocations { + out[a.LotID] = a.Amount + } + return out +} + +func TestWalletProcessor_CreditCreatesLedgerRowAndLot(t *testing.T) { + e := newWalletTestEnv(t) + c := e.customer() + + res := e.credit(t, earn(c, 100, e.at(24*time.Hour))) + + assert.Equal(t, int64(100), res.Transaction.Amount) + assert.Equal(t, int64(100), res.Transaction.BalanceAfter) + assert.Equal(t, e.org, res.Transaction.OrganizationID, "organization comes from the wallet") + require.Len(t, res.Lots, 1) + assert.Equal(t, res.Transaction.ID, res.Lots[0].SourceTransactionID) + assert.Equal(t, e.at(24*time.Hour), res.Lots[0].ExpiresAt) + assert.Equal(t, int64(100), e.balance(t, c)) + assert.Equal(t, 1, e.repo.locks[c], "credit locks the wallet itself") +} + +func TestWalletProcessor_CreditWithoutLotsMakesOneNonExpiringLot(t *testing.T) { + e := newWalletTestEnv(t) + c := e.customer() + in := earn(c, 40, nil) + in.Lots = nil + + res := e.credit(t, in) + require.Len(t, res.Lots, 1) + assert.Equal(t, int64(40), res.Lots[0].OriginalAmount) + assert.Nil(t, res.Lots[0].ExpiresAt) +} + +func TestWalletProcessor_CreditRejectsLotsThatDoNotAddUp(t *testing.T) { + e := newWalletTestEnv(t) + c := e.customer() + in := earn(c, 100, nil) + in.Lots = []WalletLotInput{{Amount: 60}, {Amount: 30}} + + _, err := e.p.Credit(e.ctx, in) + assert.ErrorIs(t, err, ErrWalletInvalidEntry) + + in.Lots = []WalletLotInput{{Amount: 100}, {Amount: 0}} + _, err = e.p.Credit(e.ctx, in) + assert.ErrorIs(t, err, ErrWalletInvalidEntry) + assert.Empty(t, e.repo.transactions) +} + +func TestWalletProcessor_DebitAcrossSeveralLots(t *testing.T) { + e := newWalletTestEnv(t) + c := e.customer() + first := e.credit(t, earn(c, 30, e.at(1*time.Hour))).Lots[0] + second := e.credit(t, earn(c, 50, e.at(2*time.Hour))).Lots[0] + third := e.credit(t, earn(c, 40, e.at(3*time.Hour))).Lots[0] + + res, err := e.p.Debit(e.ctx, pay(c, 70)) + require.NoError(t, err) + + assert.Equal(t, int64(-70), res.Transaction.Amount) + assert.Equal(t, int64(50), res.Transaction.BalanceAfter) + assert.Equal(t, map[uuid.UUID]int64{first.ID: 30, second.ID: 40}, allocationsOf(res)) + assert.Equal(t, first.ID, res.Allocations[0].LotID, "allocations are reported in the order used") + assert.Equal(t, e.at(1*time.Hour), res.Allocations[0].ExpiresAt) + + lots, _ := e.repo.GetLotsByIDs(e.ctx, []uuid.UUID{first.ID, second.ID, third.ID}) + remaining := map[uuid.UUID]int64{} + for _, l := range lots { + remaining[l.ID] = l.RemainingAmount + } + assert.Equal(t, map[uuid.UUID]int64{first.ID: 0, second.ID: 10, third.ID: 40}, remaining) +} + +// K9: soonest expiry first, lots without an expiry last and oldest first among them, +// expired lots never. +func TestWalletProcessor_DebitFollowsLotOrder(t *testing.T) { + e := newWalletTestEnv(t) + c := e.customer() + neverOld := e.credit(t, earn(c, 10, nil)).Lots[0] + late := e.credit(t, earn(c, 10, e.at(48*time.Hour))).Lots[0] + soon := e.credit(t, earn(c, 10, e.at(1*time.Hour))).Lots[0] + neverNew := e.credit(t, earn(c, 10, nil)).Lots[0] + e.credit(t, earn(c, 10, e.at(-1*time.Hour))) // already expired + + var order []uuid.UUID + for i := 0; i < 4; i++ { + res, err := e.p.Debit(e.ctx, pay(c, 10)) + require.NoError(t, err) + require.Len(t, res.Allocations, 1) + order = append(order, res.Allocations[0].LotID) + } + assert.Equal(t, []uuid.UUID{soon.ID, late.ID, neverOld.ID, neverNew.ID}, order) + + // The expired lot still counts in the balance until the expiry job removes it, + // but it cannot be spent (§7.3). + assert.Equal(t, int64(10), e.balance(t, c)) + _, err := e.p.Debit(e.ctx, pay(c, 10)) + assert.ErrorIs(t, err, repository.ErrWalletInsufficientBalance) +} + +func TestWalletProcessor_DebitOverBalanceChangesNothing(t *testing.T) { + e := newWalletTestEnv(t) + c := e.customer() + e.credit(t, earn(c, 50, nil)) + + _, err := e.p.Debit(e.ctx, pay(c, 51)) + assert.ErrorIs(t, err, repository.ErrWalletInsufficientBalance) + assert.Equal(t, int64(50), e.balance(t, c)) + assert.Len(t, e.repo.transactions, 1) + assert.Empty(t, e.repo.allocations) + + // A customer who never had a wallet has nothing to spend. + _, err = e.p.Debit(e.ctx, pay(e.customer(), 1)) + assert.ErrorIs(t, err, repository.ErrWalletInsufficientBalance) +} + +func reversal(customerID uuid.UUID, amount int64, earnRes *WalletResult) WalletDebitInput { + in := WalletDebitInput{WalletEntry: WalletEntry{ + CustomerID: customerID, + Currency: constants.WalletCurrencyPoint, + Type: constants.WalletTxTypeEarnReversal, + Amount: amount, + ReferenceType: constants.WalletRefTypeOrder, + ReferenceID: earnRes.Transaction.ReferenceID, + ReversesTransactionID: &earnRes.Transaction.ID, + OutletID: earnRes.Transaction.OutletID, + Description: "Batal", + }} + for _, lot := range earnRes.Lots { + in.PreferredLotIDs = append(in.PreferredLotIDs, lot.ID) + } + return in +} + +func TestWalletProcessor_DebitUpToWithShortfall(t *testing.T) { + e := newWalletTestEnv(t) + c := e.customer() + earned := e.credit(t, earn(c, 100, nil)) + _, err := e.p.Debit(e.ctx, pay(c, 70)) + require.NoError(t, err) + + res, err := e.p.DebitUpTo(e.ctx, reversal(c, 100, earned)) + require.NoError(t, err) + assert.Equal(t, int64(-30), res.Transaction.Amount) + assert.Equal(t, int64(70), res.Shortfall) + assert.Equal(t, int64(100), res.Transaction.Metadata["requested_amount"]) + assert.Equal(t, int64(70), res.Transaction.Metadata["shortfall"]) + assert.Equal(t, int64(0), e.balance(t, c)) + + // Nothing left: no ledger row, the whole amount is shortfall. + in := reversal(c, 5, earned) + res, err = e.p.DebitUpTo(e.ctx, in) + require.NoError(t, err) + assert.Nil(t, res.Transaction) + assert.Equal(t, int64(5), res.Shortfall) + assert.Len(t, e.repo.transactions, 3) +} + +// A reversal draws from the lots its EARN created first (F10), even when an older lot +// would come first in K9 order, and even when that lot has expired. +func TestWalletProcessor_DebitDrawsPreferredLotsFirst(t *testing.T) { + e := newWalletTestEnv(t) + c := e.customer() + older := e.credit(t, earn(c, 50, e.at(1*time.Hour))).Lots[0] + earned := e.credit(t, earn(c, 20, e.at(-1*time.Hour))) + + res, err := e.p.Debit(e.ctx, reversal(c, 30, earned)) + require.NoError(t, err) + require.Len(t, res.Allocations, 2) + assert.Equal(t, earned.Lots[0].ID, res.Allocations[0].LotID) + assert.Equal(t, int64(20), res.Allocations[0].Amount) + assert.Equal(t, older.ID, res.Allocations[1].LotID) + assert.Equal(t, int64(10), res.Allocations[1].Amount) +} + +func TestWalletProcessor_DebitRejectsSomeoneElsesLot(t *testing.T) { + e := newWalletTestEnv(t) + a, b := e.customer(), e.customer() + e.credit(t, earn(a, 10, nil)) + other := e.credit(t, earn(b, 10, nil)) + + in := pay(a, 5) + in.PreferredLotIDs = []uuid.UUID{other.Lots[0].ID} + _, err := e.p.Debit(e.ctx, in) + assert.ErrorIs(t, err, ErrWalletInvalidEntry) + + in.PreferredLotIDs = []uuid.UUID{uuid.New()} + _, err = e.p.Debit(e.ctx, in) + assert.ErrorIs(t, err, ErrWalletInvalidEntry) +} + +func TestWalletProcessor_ExpireDrawsTheExpiredLot(t *testing.T) { + e := newWalletTestEnv(t) + c := e.customer() + e.credit(t, earn(c, 10, nil)) + expired := e.credit(t, earn(c, 25, e.at(-1*time.Hour))).Lots[0] + + res, err := e.p.Debit(e.ctx, WalletDebitInput{ + WalletEntry: WalletEntry{ + CustomerID: c, + Currency: constants.WalletCurrencyPoint, + Type: constants.WalletTxTypeExpire, + Amount: expired.RemainingAmount, + ReferenceType: constants.WalletRefTypeLot, + ReferenceID: expired.ID, + Description: "Kedaluwarsa", + IdempotencyKey: "expire:" + expired.ID.String(), + }, + PreferredLotIDs: []uuid.UUID{expired.ID}, + }) + require.NoError(t, err) + assert.Equal(t, map[uuid.UUID]int64{expired.ID: 25}, allocationsOf(res)) + assert.Equal(t, int64(10), e.balance(t, c)) +} + +func TestWalletProcessor_IdempotentCredit(t *testing.T) { + e := newWalletTestEnv(t) + c := e.customer() + in := earn(c, 100, nil) + in.IdempotencyKey = "earn:order-1" + + first := e.credit(t, in) + in.ReferenceID = first.Transaction.ReferenceID + second := e.credit(t, in) + + assert.True(t, second.Replayed) + assert.False(t, first.Replayed) + assert.Equal(t, first.Transaction.ID, second.Transaction.ID) + assert.Equal(t, first.Lots[0].ID, second.Lots[0].ID) + assert.Equal(t, int64(100), e.balance(t, c)) + assert.Len(t, e.repo.transactions, 1) +} + +func TestWalletProcessor_IdempotentDebit(t *testing.T) { + e := newWalletTestEnv(t) + c := e.customer() + e.credit(t, earn(c, 30, e.at(time.Hour))) + e.credit(t, earn(c, 30, nil)) + in := pay(c, 40) + in.IdempotencyKey = "pay:1" + + first, err := e.p.Debit(e.ctx, in) + require.NoError(t, err) + second, err := e.p.Debit(e.ctx, in) + require.NoError(t, err) + + assert.True(t, second.Replayed) + assert.Equal(t, first.Transaction.ID, second.Transaction.ID) + assert.ElementsMatch(t, first.Allocations, second.Allocations) + assert.Equal(t, int64(20), e.balance(t, c)) +} + +func TestWalletProcessor_IdempotentDebitUpToKeepsShortfall(t *testing.T) { + e := newWalletTestEnv(t) + c := e.customer() + earned := e.credit(t, earn(c, 100, nil)) + _, err := e.p.Debit(e.ctx, pay(c, 60)) + require.NoError(t, err) + in := reversal(c, 100, earned) + in.IdempotencyKey = "reverse:order-1" + + first, err := e.p.DebitUpTo(e.ctx, in) + require.NoError(t, err) + e.credit(t, earn(c, 500, nil)) // new balance must not be taken by the retry + second, err := e.p.DebitUpTo(e.ctx, in) + require.NoError(t, err) + + assert.True(t, second.Replayed) + assert.Equal(t, first.Transaction.ID, second.Transaction.ID) + assert.Equal(t, int64(60), second.Shortfall) + assert.Equal(t, int64(500), e.balance(t, c)) +} + +func TestWalletProcessor_IdempotencyKeyReusedForAnotherOperation(t *testing.T) { + e := newWalletTestEnv(t) + c := e.customer() + in := earn(c, 100, nil) + in.IdempotencyKey = "k" + e.credit(t, in) + + other := earn(c, 99, nil) + other.IdempotencyKey = "k" + _, err := e.p.Credit(e.ctx, other) + assert.ErrorIs(t, err, ErrWalletIdempotencyConflict) + + debit := pay(c, 100) + debit.IdempotencyKey = "k" + _, err = e.p.Debit(e.ctx, debit) + assert.ErrorIs(t, err, ErrWalletIdempotencyConflict) + + otherCustomer := earn(e.customer(), 100, nil) + otherCustomer.IdempotencyKey = "k" + _, err = e.p.Credit(e.ctx, otherCustomer) + assert.ErrorIs(t, err, ErrWalletIdempotencyConflict) +} + +// A transfer debits the sender and credits the receiver with lots that keep the +// sender's expiry (K9), following the example in §8. +func TestWalletProcessor_TransferCarriesExpiry(t *testing.T) { + e := newWalletTestEnv(t) + a, b := e.customer(), e.customer() + dec := e.credit(t, earn(a, 100, e.at(30*24*time.Hour))).Lots[0] + jan := e.credit(t, earn(a, 50, e.at(60*24*time.Hour))).Lots[0] + + require.NoError(t, e.p.LockWallets(e.ctx, a, b)) + group, outID, inID := uuid.New(), uuid.New(), uuid.New() + out, err := e.p.Debit(e.ctx, WalletDebitInput{WalletEntry: WalletEntry{ + TransactionID: outID, CustomerID: a, Currency: constants.WalletCurrencyPoint, + Type: constants.WalletTxTypeTransferOut, Amount: 120, + ReferenceType: constants.WalletRefTypeWalletTx, ReferenceID: inID, + GroupID: &group, CounterpartyCustomerID: &b, Description: "Transfer ke B", + }}) + require.NoError(t, err) + received, err := e.p.Credit(e.ctx, WalletCreditInput{ + WalletEntry: WalletEntry{ + TransactionID: inID, CustomerID: b, Currency: constants.WalletCurrencyPoint, + Type: constants.WalletTxTypeTransferIn, Amount: 120, + ReferenceType: constants.WalletRefTypeWalletTx, ReferenceID: outID, + GroupID: &group, CounterpartyCustomerID: &a, Description: "Transfer dari A", + }, + Lots: out.CarryOver(), + }) + require.NoError(t, err) + + assert.Equal(t, outID, out.Transaction.ID) + assert.Equal(t, inID, received.Transaction.ID) + require.Len(t, received.Lots, 2) + assert.Equal(t, int64(100), received.Lots[0].OriginalAmount) + assert.Equal(t, dec.ExpiresAt, received.Lots[0].ExpiresAt) + assert.Equal(t, &dec.ID, received.Lots[0].OriginLotID) + assert.Equal(t, int64(20), received.Lots[1].OriginalAmount) + assert.Equal(t, jan.ExpiresAt, received.Lots[1].ExpiresAt) + assert.Equal(t, &jan.ID, received.Lots[1].OriginLotID) + assert.Equal(t, int64(30), e.balance(t, a)) + assert.Equal(t, int64(120), e.balance(t, b)) +} + +func TestWalletProcessor_RejectsEntriesThatBreakTheTypeRules(t *testing.T) { + c := uuid.New() + outlet := ptr(uuid.New()) + + credits := map[string]func(*WalletCreditInput){ + "unknown type": func(in *WalletCreditInput) { in.Type = "BONUS" }, + "debit-only type as credit": func(in *WalletCreditInput) { + in.Type = constants.WalletTxTypePayment + in.ReferenceType = constants.WalletRefTypePayment + }, + "unknown currency": func(in *WalletCreditInput) { in.Currency = "GOLD" }, + "zero amount": func(in *WalletCreditInput) { in.Amount = 0; in.Lots = nil }, + "negative amount": func(in *WalletCreditInput) { in.Amount = -5; in.Lots = nil }, + "wrong reference type": func(in *WalletCreditInput) { in.ReferenceType = constants.WalletRefTypeUser }, + "missing reference id": func(in *WalletCreditInput) { in.ReferenceID = uuid.Nil }, + "missing description": func(in *WalletCreditInput) { in.Description = " " }, + "EARN without outlet": func(in *WalletCreditInput) { in.OutletID = nil }, + "missing customer": func(in *WalletCreditInput) { in.CustomerID = uuid.Nil }, + "EXCHANGE_IN in COIN": func(in *WalletCreditInput) { + in.Type = constants.WalletTxTypeExchangeIn + in.Currency = constants.WalletCurrencyCoin + in.ReferenceType = constants.WalletRefTypeWalletTx + in.GroupID = ptr(uuid.New()) + }, + "TRANSFER_IN to self": func(in *WalletCreditInput) { + in.Type = constants.WalletTxTypeTransferIn + in.ReferenceType = constants.WalletRefTypeWalletTx + in.GroupID = ptr(uuid.New()) + in.CounterpartyCustomerID = &in.CustomerID + }, + "TRANSFER_IN without group": func(in *WalletCreditInput) { + in.Type = constants.WalletTxTypeTransferIn + in.ReferenceType = constants.WalletRefTypeWalletTx + in.CounterpartyCustomerID = ptr(uuid.New()) + }, + "PAYMENT_REFUND without source": func(in *WalletCreditInput) { + in.Type = constants.WalletTxTypePaymentRefund + in.ReferenceType = constants.WalletRefTypePayment + }, + "ADJUSTMENT without reason": func(in *WalletCreditInput) { + in.Type = constants.WalletTxTypeAdjustment + in.ReferenceType = constants.WalletRefTypeUser + in.CreatedByUser = ptr(uuid.New()) + }, + "ADJUSTMENT blank reason": func(in *WalletCreditInput) { + in.Type = constants.WalletTxTypeAdjustment + in.ReferenceType = constants.WalletRefTypeUser + in.CreatedByUser = ptr(uuid.New()) + in.Reason = ptr(" ") + }, + "ADJUSTMENT without admin": func(in *WalletCreditInput) { + in.Type = constants.WalletTxTypeAdjustment + in.ReferenceType = constants.WalletRefTypeUser + in.Reason = ptr("komplain") + }, + "MIGRATION wrong reference": func(in *WalletCreditInput) { + in.Type = constants.WalletTxTypeMigration + in.ReferenceType = constants.WalletRefTypeOrder + }, + } + for name, mutate := range credits { + t.Run("credit/"+name, func(t *testing.T) { + e := newWalletTestEnv(t) + e.repo.customers[c] = e.org + in := earn(c, 10, nil) + in.OutletID = outlet + mutate(&in) + _, err := e.p.Credit(e.ctx, in) + assert.ErrorIs(t, err, ErrWalletInvalidEntry) + assert.Empty(t, e.repo.transactions) + }) + } + + debits := map[string]func(*WalletDebitInput){ + "credit-only type as debit": func(in *WalletDebitInput) { + in.Type = constants.WalletTxTypeMigration + in.ReferenceType = constants.WalletRefTypeLegacyPoints + }, + "PAYMENT in COIN": func(in *WalletDebitInput) { in.Currency = constants.WalletCurrencyCoin }, + "GAME_SPEND in POINT": func(in *WalletDebitInput) { + in.Type = constants.WalletTxTypeGameSpend + in.ReferenceType = constants.WalletRefTypeGamePlay + }, + "EXPIRE not pointing at a lot": func(in *WalletDebitInput) { + in.Type = constants.WalletTxTypeExpire + in.ReferenceType = constants.WalletRefTypeOrder + }, + "EARN_REVERSAL without source": func(in *WalletDebitInput) { + in.Type = constants.WalletTxTypeEarnReversal + in.ReferenceType = constants.WalletRefTypeOrder + }, + "TRANSFER_OUT without counterparty": func(in *WalletDebitInput) { + in.Type = constants.WalletTxTypeTransferOut + in.ReferenceType = constants.WalletRefTypeWalletTx + in.GroupID = ptr(uuid.New()) + }, + "REWARD_REDEEM wrong reference": func(in *WalletDebitInput) { + in.Type = constants.WalletTxTypeRewardRedeem + in.ReferenceType = constants.WalletRefTypeOrder + }, + } + for name, mutate := range debits { + t.Run("debit/"+name, func(t *testing.T) { + e := newWalletTestEnv(t) + e.repo.customers[c] = e.org + e.credit(t, earn(c, 100, nil)) + in := pay(c, 10) + mutate(&in) + _, err := e.p.Debit(e.ctx, in) + assert.ErrorIs(t, err, ErrWalletInvalidEntry) + assert.Len(t, e.repo.transactions, 1) + }) + } +} + +func TestWalletProcessor_UnknownCustomer(t *testing.T) { + e := newWalletTestEnv(t) + _, err := e.p.Credit(e.ctx, earn(uuid.New(), 10, nil)) + assert.ErrorIs(t, err, repository.ErrWalletNotFound) +} diff --git a/internal/repository/wallet_repository.go b/internal/repository/wallet_repository.go index b33f761..a88e269 100644 --- a/internal/repository/wallet_repository.go +++ b/internal/repository/wallet_repository.go @@ -52,6 +52,11 @@ type WalletRepository interface { 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. @@ -204,6 +209,32 @@ func (r *walletRepository) CreateLot(ctx context.Context, lot *entities.WalletLo 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.