package repository import ( "context" "errors" "fmt" "time" "github.com/google/uuid" "gorm.io/gorm" "gorm.io/gorm/clause" "apskel-pos-be/internal/constants" "apskel-pos-be/internal/entities" ) // ErrVoucherNotFound means no voucher with that id in the organization. var ErrVoucherNotFound = errors.New("enakgame: voucher not found") // voucherCodeImportBatch is how many codes one insert writes. const voucherCodeImportBatch = 1000 // VoucherFilter selects an organization's vouchers. type VoucherFilter struct { OrganizationID uuid.UUID // Empty for every status. Statuses []string Search string Offset int Limit int } // VoucherCodeImport is one code to add to a pool. type VoucherCodeImport struct { Code string ExpiresAt *time.Time } // CatalogVoucher is a voucher a customer can redeem now, with how many are left. Nil // Available means no counted stock. type CatalogVoucher struct { entities.Voucher Available *int64 } // VoucherRepository stores EnakGame vouchers and their codes (docs/rfc-enakgame.md // ยง5.7), always scoped to an organization. type VoucherRepository interface { CreateVoucher(ctx context.Context, voucher *entities.Voucher) error GetVoucher(ctx context.Context, organizationID, id uuid.UUID) (*entities.Voucher, error) // LockVoucher is GetVoucher with the row locked until the transaction ends. LockVoucher(ctx context.Context, organizationID, id uuid.UUID) (*entities.Voucher, error) ListVouchers(ctx context.Context, filter VoucherFilter) ([]entities.Voucher, int64, error) // UpdateVoucher stores everything but the organization, the stock mode and the // status. UpdateVoucher(ctx context.Context, voucher *entities.Voucher) error SetVoucherStatus(ctx context.Context, organizationID, id uuid.UUID, status string) error // TakeStock takes one from a STATIC voucher's stock, and reports false when none is // left. TakeStock(ctx context.Context, voucherID uuid.UUID) (bool, error) // ImportCodes adds codes to a pool and returns those it added; a code the pool // already holds is skipped. ImportCodes(ctx context.Context, voucherID uuid.UUID, codes []VoucherCodeImport) ([]string, error) // CountCodes returns how many codes of a pool are in each status. CountCodes(ctx context.Context, voucherID uuid.UUID) (map[string]int64, error) // ListCodes returns a page of a pool's codes, oldest first, and the total. ListCodes(ctx context.Context, voucherID uuid.UUID, status string, offset, limit int) ([]entities.VoucherCode, int64, error) GetCode(ctx context.Context, id uuid.UUID) (*entities.VoucherCode, error) // ClaimCode gives the oldest available, unexpired code of a pool to a redemption, // skipping codes other redemptions hold at the moment, so two redemptions at once // get different codes. It returns nil when none is left. ClaimCode(ctx context.Context, voucherID, redemptionID uuid.UUID, now time.Time) (*entities.VoucherCode, error) // ExpireCodes moves at most limit available codes past their expiry to EXPIRED // and returns how many. ExpireCodes(ctx context.Context, now time.Time, limit int) (int64, error) // ListCatalog returns the vouchers a customer of the organization can redeem now: // ACTIVE, within their dates, in a stock mode that can be redeemed. ListCatalog(ctx context.Context, organizationID uuid.UUID, now time.Time, stockModes []string) ([]CatalogVoucher, error) } type voucherRepository struct { db *gorm.DB } func NewVoucherRepository(db *gorm.DB) VoucherRepository { return &voucherRepository{db: db} } func (r *voucherRepository) CreateVoucher(ctx context.Context, voucher *entities.Voucher) error { if err := DBFromContext(ctx, r.db).WithContext(ctx).Create(voucher).Error; err != nil { return fmt.Errorf("failed to create voucher: %w", err) } return nil } func (r *voucherRepository) getVoucher(ctx context.Context, organizationID, id uuid.UUID, lock bool) (*entities.Voucher, error) { q := DBFromContext(ctx, r.db).WithContext(ctx).Where("organization_id = ? AND id = ?", organizationID, id) if lock { q = q.Clauses(clause.Locking{Strength: "UPDATE"}) } var voucher entities.Voucher if err := q.First(&voucher).Error; err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return nil, ErrVoucherNotFound } return nil, fmt.Errorf("failed to read voucher: %w", err) } return &voucher, nil } func (r *voucherRepository) GetVoucher(ctx context.Context, organizationID, id uuid.UUID) (*entities.Voucher, error) { return r.getVoucher(ctx, organizationID, id, false) } func (r *voucherRepository) LockVoucher(ctx context.Context, organizationID, id uuid.UUID) (*entities.Voucher, error) { return r.getVoucher(ctx, organizationID, id, true) } func (r *voucherRepository) ListVouchers(ctx context.Context, filter VoucherFilter) ([]entities.Voucher, int64, error) { q := DBFromContext(ctx, r.db).WithContext(ctx).Model(&entities.Voucher{}).Where("organization_id = ?", filter.OrganizationID) if len(filter.Statuses) > 0 { q = q.Where("status IN ?", filter.Statuses) } if filter.Search != "" { q = q.Where("name ILIKE ?", "%"+filter.Search+"%") } var total int64 if err := q.Count(&total).Error; err != nil { return nil, 0, fmt.Errorf("failed to count vouchers: %w", err) } var vouchers []entities.Voucher if err := q.Order("created_at DESC, id").Offset(filter.Offset).Limit(filter.Limit).Find(&vouchers).Error; err != nil { return nil, 0, fmt.Errorf("failed to list vouchers: %w", err) } return vouchers, total, nil } func (r *voucherRepository) UpdateVoucher(ctx context.Context, v *entities.Voucher) error { terms := v.Terms if len(terms) == 0 { terms = entities.JSONDocument(`{}`) } result := DBFromContext(ctx, r.db).WithContext(ctx).Exec(` UPDATE vouchers SET name = ?, description = ?, image_url = ?, voucher_type = ?, face_value = ?, point_cost = ?, business_cost = ?, stock = ?, provider = ?, provider_ref = ?, max_per_customer = ?, valid_from = ?, valid_until = ?, terms = ?::jsonb, updated_at = NOW() WHERE organization_id = ? AND id = ?`, v.Name, v.Description, v.ImageURL, v.VoucherType, v.FaceValue, v.PointCost, v.BusinessCost, v.Stock, v.Provider, v.ProviderRef, v.MaxPerCustomer, v.ValidFrom, v.ValidUntil, terms, v.OrganizationID, v.ID) if result.Error != nil { return fmt.Errorf("failed to update voucher: %w", result.Error) } if result.RowsAffected == 0 { return ErrVoucherNotFound } return nil } func (r *voucherRepository) SetVoucherStatus(ctx context.Context, organizationID, id uuid.UUID, status string) error { result := DBFromContext(ctx, r.db).WithContext(ctx).Exec(` UPDATE vouchers SET status = ?, updated_at = NOW() WHERE organization_id = ? AND id = ?`, status, organizationID, id) if result.Error != nil { return fmt.Errorf("failed to change voucher status: %w", result.Error) } if result.RowsAffected == 0 { return ErrVoucherNotFound } return nil } func (r *voucherRepository) TakeStock(ctx context.Context, voucherID uuid.UUID) (bool, error) { result := DBFromContext(ctx, r.db).WithContext(ctx).Exec(` UPDATE vouchers SET stock = stock - 1, updated_at = NOW() WHERE id = ? AND stock > 0`, voucherID) if result.Error != nil { return false, fmt.Errorf("failed to take voucher stock: %w", result.Error) } return result.RowsAffected == 1, nil } func (r *voucherRepository) ImportCodes(ctx context.Context, voucherID uuid.UUID, codes []VoucherCodeImport) ([]string, error) { db := DBFromContext(ctx, r.db).WithContext(ctx) var added []string for start := 0; start < len(codes); start += voucherCodeImportBatch { end := start + voucherCodeImportBatch if end > len(codes) { end = len(codes) } rows := make([]entities.VoucherCode, 0, end-start) for _, c := range codes[start:end] { rows = append(rows, entities.VoucherCode{ ID: uuid.New(), VoucherID: voucherID, Code: c.Code, Status: constants.VoucherCodeAvailable, ExpiresAt: c.ExpiresAt, }) } err := db.Clauses(clause.OnConflict{Columns: []clause.Column{{Name: "voucher_id"}, {Name: "code"}}, DoNothing: true}). Create(&rows).Error if err != nil { return nil, fmt.Errorf("failed to import voucher codes: %w", err) } // A skipped code kept the row of the pool's earlier copy, so only the new ids // are in the table. ids := make([]uuid.UUID, 0, len(rows)) for _, row := range rows { ids = append(ids, row.ID) } var created []string if err := db.Model(&entities.VoucherCode{}).Where("id IN ?", ids).Order("code").Pluck("code", &created).Error; err != nil { return nil, fmt.Errorf("failed to read imported voucher codes: %w", err) } added = append(added, created...) } return added, nil } func (r *voucherRepository) CountCodes(ctx context.Context, voucherID uuid.UUID) (map[string]int64, error) { var rows []struct { Status string Count int64 } err := DBFromContext(ctx, r.db).WithContext(ctx).Raw(` SELECT status, COUNT(*) AS count FROM voucher_codes WHERE voucher_id = ? GROUP BY status`, voucherID).Scan(&rows).Error if err != nil { return nil, fmt.Errorf("failed to count voucher codes: %w", err) } counts := map[string]int64{} for _, row := range rows { counts[row.Status] = row.Count } return counts, nil } func (r *voucherRepository) ListCodes(ctx context.Context, voucherID uuid.UUID, status string, offset, limit int) ([]entities.VoucherCode, int64, error) { q := DBFromContext(ctx, r.db).WithContext(ctx).Model(&entities.VoucherCode{}).Where("voucher_id = ?", voucherID) if status != "" { q = q.Where("status = ?", status) } var total int64 if err := q.Count(&total).Error; err != nil { return nil, 0, fmt.Errorf("failed to count voucher codes: %w", err) } var codes []entities.VoucherCode if err := q.Order("created_at, code").Offset(offset).Limit(limit).Find(&codes).Error; err != nil { return nil, 0, fmt.Errorf("failed to list voucher codes: %w", err) } return codes, total, nil } func (r *voucherRepository) GetCode(ctx context.Context, id uuid.UUID) (*entities.VoucherCode, error) { var code entities.VoucherCode if err := DBFromContext(ctx, r.db).WithContext(ctx).Where("id = ?", id).First(&code).Error; err != nil { return nil, fmt.Errorf("failed to read voucher code: %w", err) } return &code, nil } func (r *voucherRepository) ClaimCode(ctx context.Context, voucherID, redemptionID uuid.UUID, now time.Time) (*entities.VoucherCode, error) { var claimed []entities.VoucherCode err := DBFromContext(ctx, r.db).WithContext(ctx).Raw(` UPDATE voucher_codes SET status = ?, redemption_id = ?, updated_at = NOW() WHERE id = ( SELECT id FROM voucher_codes WHERE voucher_id = ? AND status = ? AND (expires_at IS NULL OR expires_at > ?) ORDER BY created_at, id LIMIT 1 FOR UPDATE SKIP LOCKED) RETURNING *`, constants.VoucherCodeRedeemed, redemptionID, voucherID, constants.VoucherCodeAvailable, now).Scan(&claimed).Error if err != nil { return nil, fmt.Errorf("failed to claim voucher code: %w", err) } if len(claimed) == 0 { return nil, nil } return &claimed[0], nil } func (r *voucherRepository) ExpireCodes(ctx context.Context, now time.Time, limit int) (int64, error) { result := DBFromContext(ctx, r.db).WithContext(ctx).Exec(` UPDATE voucher_codes SET status = ?, updated_at = NOW() WHERE id IN ( SELECT id FROM voucher_codes WHERE status = ? AND expires_at IS NOT NULL AND expires_at <= ? ORDER BY expires_at LIMIT ? FOR UPDATE SKIP LOCKED)`, constants.VoucherCodeExpired, constants.VoucherCodeAvailable, now, limit) if result.Error != nil { return 0, fmt.Errorf("failed to expire voucher codes: %w", result.Error) } return result.RowsAffected, nil } func (r *voucherRepository) ListCatalog(ctx context.Context, organizationID uuid.UUID, now time.Time, stockModes []string) ([]CatalogVoucher, error) { db := DBFromContext(ctx, r.db).WithContext(ctx) var vouchers []entities.Voucher err := db.Where(`organization_id = ? AND status = ? AND stock_mode IN ? AND (valid_from IS NULL OR valid_from <= ?) AND (valid_until IS NULL OR valid_until > ?)`, organizationID, constants.VoucherStatusActive, stockModes, now, now). Order("point_cost, name, id").Find(&vouchers).Error if err != nil { return nil, fmt.Errorf("failed to list voucher catalog: %w", err) } var pools []uuid.UUID for _, v := range vouchers { if v.StockMode == constants.VoucherStockCodePool { pools = append(pools, v.ID) } } available := map[uuid.UUID]int64{} if len(pools) > 0 { var rows []struct { VoucherID string Count int64 } err := db.Raw(` SELECT voucher_id::text AS voucher_id, COUNT(*) AS count FROM voucher_codes WHERE voucher_id IN ? AND status = ? AND (expires_at IS NULL OR expires_at > ?) GROUP BY voucher_id`, pools, constants.VoucherCodeAvailable, now).Scan(&rows).Error if err != nil { return nil, fmt.Errorf("failed to count available voucher codes: %w", err) } for _, row := range rows { id, _ := uuid.Parse(row.VoucherID) available[id] = row.Count } } out := make([]CatalogVoucher, 0, len(vouchers)) for _, v := range vouchers { c := CatalogVoucher{Voucher: v} switch v.StockMode { case constants.VoucherStockStatic: c.Available = v.Stock case constants.VoucherStockCodePool: n := available[v.ID] c.Available = &n } out = append(out, c) } return out, nil }