79 lines
2.3 KiB
Go
79 lines
2.3 KiB
Go
package reward
|
|
|
|
import (
|
|
"context"
|
|
"database/sql"
|
|
"fmt"
|
|
"testing"
|
|
"time"
|
|
|
|
"github.com/wangjia/pangolin/server/internal/notices"
|
|
)
|
|
|
|
// failingGranter 总是拒绝发放,用于验证「发放失败 → 通知(与本事务其余写入)一并回滚」。
|
|
type failingGranter struct{}
|
|
|
|
func (failingGranter) GrantRewardTx(ctx context.Context, tx *sql.Tx, userID int64, days int, source, auditAction, ref string) (int64, time.Time, error) {
|
|
return 0, time.Time{}, fmt.Errorf("boom")
|
|
}
|
|
|
|
func TestOnRegister_InsertsRewardNotices(t *testing.T) {
|
|
db := openDB(t)
|
|
seedU(t, db, 1, "inviter")
|
|
seedU(t, db, 2, "invitee")
|
|
s := newSvc(t, db)
|
|
s.SetNoticer(notices.NewStore(db))
|
|
code, _ := s.EnsureCode(context.Background(), 1)
|
|
s.OnRegister(context.Background(), 2, code, "dev-2")
|
|
var n int
|
|
db.QueryRow(`SELECT COUNT(*) FROM notices WHERE type='reward' AND user_id IN (1,2)`).Scan(&n)
|
|
if n != 2 {
|
|
t.Fatalf("注册段应双方各一条到账通知, got %d", n)
|
|
}
|
|
}
|
|
|
|
func TestClaimTelegram_NoticeRollsBackWithGrantFailure(t *testing.T) {
|
|
db := openDB(t)
|
|
seedU(t, db, 1, "u1")
|
|
s := NewService(db, NewStore(db), failingGranter{}, Config{TgDays: 3}, time.Now)
|
|
s.SetTelegram("bot", "@ch", "tok", "sec")
|
|
s.SetMemberChecker(fakeChecker{member: true})
|
|
s.SetNoticer(notices.NewStore(db))
|
|
|
|
if _, err := s.ClaimTelegram(context.Background(), 1, "tg-1"); err == nil {
|
|
t.Fatalf("expected grant failure error, got nil")
|
|
}
|
|
var n int
|
|
db.QueryRow(`SELECT COUNT(*) FROM notices`).Scan(&n)
|
|
if n != 0 {
|
|
t.Fatalf("grant 失败应整事务回滚,notices should be 0, got %d", n)
|
|
}
|
|
}
|
|
|
|
func TestOnFirstPaidTx_InsertsRewardNotices(t *testing.T) {
|
|
db := openDB(t)
|
|
seedU(t, db, 1, "inviter")
|
|
seedU(t, db, 2, "invitee")
|
|
s := newSvc(t, db)
|
|
s.SetNoticer(notices.NewStore(db))
|
|
code, _ := s.EnsureCode(context.Background(), 1)
|
|
s.OnRegister(context.Background(), 2, code, "dev-2")
|
|
|
|
tx, err := db.Begin()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := s.OnFirstPaidTx(context.Background(), tx, 2, time.Now().UTC()); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := tx.Commit(); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
var n int
|
|
db.QueryRow(`SELECT COUNT(*) FROM notices WHERE type='reward' AND user_id IN (1,2)`).Scan(&n)
|
|
// 注册段 2 条 + 首充段 2 条 = 4 条
|
|
if n != 4 {
|
|
t.Fatalf("注册+首充共应 4 条到账通知, got %d", n)
|
|
}
|
|
}
|