diff --git a/server/internal/reward/store.go b/server/internal/reward/store.go new file mode 100644 index 0000000..605c8ea --- /dev/null +++ b/server/internal/reward/store.go @@ -0,0 +1,156 @@ +// Package reward 承载邀请奖励与奖励任务的数据访问 + 发奖 + 防刷。 +package reward + +import ( + "context" + "database/sql" + "errors" + "strings" + "time" +) + +var ErrClaimExists = errors.New("reward: claim already exists") + +type Store struct{ db *sql.DB } + +func NewStore(db *sql.DB) *Store { return &Store{db: db} } + +// EnsureInviteCode 惰性生成邀请码;已有则返回旧值。gen 生成候选码(冲突时重试到成功)。 +func (s *Store) EnsureInviteCode(ctx context.Context, userID int64, gen func() string) (string, error) { + var existing sql.NullString + if err := s.db.QueryRowContext(ctx, `SELECT invite_code FROM users WHERE id=?`, userID).Scan(&existing); err != nil { + return "", err + } + if existing.Valid && existing.String != "" { + return existing.String, nil + } + for i := 0; i < 5; i++ { + code := gen() + _, err := s.db.ExecContext(ctx, `UPDATE users SET invite_code=? WHERE id=? AND invite_code IS NULL`, code, userID) + if err != nil { + if isDup(err) { + continue + } + return "", err + } + // 读回(并发下可能是别的并发写入的值) + var got sql.NullString + if err := s.db.QueryRowContext(ctx, `SELECT invite_code FROM users WHERE id=?`, userID).Scan(&got); err != nil { + return "", err + } + if got.Valid && got.String != "" { + return got.String, nil + } + } + return "", errors.New("reward: invite code generation exhausted") +} + +func (s *Store) ResolveInviteCode(ctx context.Context, code string) (int64, bool, error) { + var id int64 + err := s.db.QueryRowContext(ctx, `SELECT id FROM users WHERE invite_code=? AND status='active'`, code).Scan(&id) + if err == sql.ErrNoRows { + return 0, false, nil + } + return id, err == nil, err +} + +func (s *Store) DeviceUsedByOther(ctx context.Context, deviceUUID string, exceptUserID int64) (bool, error) { + if deviceUUID == "" { + return false, nil + } + var n int + err := s.db.QueryRowContext(ctx, + `SELECT COUNT(*) FROM devices WHERE uuid=? AND user_id<>?`, deviceUUID, exceptUserID).Scan(&n) + return n > 0, err +} + +func (s *Store) RegRewardCountThisMonth(ctx context.Context, inviterID int64, since time.Time) (int, error) { + var n int + err := s.db.QueryRowContext(ctx, + `SELECT COUNT(*) FROM referrals WHERE inviter_id=? AND status='reg_rewarded' AND reg_rewarded_at>=?`, + inviterID, since).Scan(&n) + return n, err +} + +func (s *Store) InsertReferralTx(ctx context.Context, tx *sql.Tx, inviterID, inviteeID int64, + deviceUUID, status string, regRewardedAt *time.Time, now time.Time) error { + _, err := tx.ExecContext(ctx, + `INSERT INTO referrals (inviter_id,invitee_id,device_uuid,status,reg_rewarded_at,created_at) + VALUES (?,?,?,?,?,?)`, inviterID, inviteeID, deviceUUID, status, regRewardedAt, now) + if isDup(err) { + return ErrClaimExists + } + return err +} + +func (s *Store) ReferralByInvitee(ctx context.Context, tx *sql.Tx, inviteeID int64) (int64, string, bool, error) { + var inviter int64 + var status string + err := tx.QueryRowContext(ctx, `SELECT inviter_id,status FROM referrals WHERE invitee_id=?`, inviteeID). + Scan(&inviter, &status) + if err == sql.ErrNoRows { + return 0, "", false, nil + } + return inviter, status, err == nil, err +} + +func (s *Store) MarkFirstPaidTx(ctx context.Context, tx *sql.Tx, userID int64, at time.Time) (bool, error) { + res, err := tx.ExecContext(ctx, + `UPDATE users SET first_paid_at=? WHERE id=? AND first_paid_at IS NULL`, at, userID) + if err != nil { + return false, err + } + n, _ := res.RowsAffected() + return n == 1, nil +} + +func (s *Store) MarkPaidRewardedTx(ctx context.Context, tx *sql.Tx, inviteeID int64, at time.Time) error { + _, err := tx.ExecContext(ctx, + `UPDATE referrals SET status='paid_rewarded', paid_rewarded_at=? WHERE invitee_id=?`, at, inviteeID) + return err +} + +func (s *Store) InsertClaimTx(ctx context.Context, tx *sql.Tx, userID int64, taskKey, externalRef string, + days int, now time.Time) error { + _, err := tx.ExecContext(ctx, + `INSERT INTO reward_claims (user_id,task_key,external_ref,granted_days,granted_at) + VALUES (?,?,?,?,?)`, userID, taskKey, externalRef, days, now) + if isDup(err) { + return ErrClaimExists + } + return err +} + +func (s *Store) TelegramClaimed(ctx context.Context, userID int64) (bool, error) { + var n int + err := s.db.QueryRowContext(ctx, + `SELECT COUNT(*) FROM reward_claims WHERE user_id=? AND task_key='telegram_join'`, userID).Scan(&n) + return n > 0, err +} + +// Summary: invited=已绑定人数;converted=已首充奖励人数;earnedDays=本人从奖励得到的总天数(audit_log 累加)。 +func (s *Store) Summary(ctx context.Context, userID int64) (invited, converted, earnedDays int, err error) { + if err = s.db.QueryRowContext(ctx, `SELECT COUNT(*) FROM referrals WHERE inviter_id=?`, userID).Scan(&invited); err != nil { + return + } + if err = s.db.QueryRowContext(ctx, + `SELECT COUNT(*) FROM referrals WHERE inviter_id=? AND status='paid_rewarded'`, userID).Scan(&converted); err != nil { + return + } + // earnedDays: audit_log 里 actor=user: 且 action∈奖励动作,meta.days 累加(简化:reward_claims + referrals 估算) + var tg, reg, paid int + s.db.QueryRowContext(ctx, `SELECT COALESCE(SUM(granted_days),0) FROM reward_claims WHERE user_id=?`, userID).Scan(&tg) + s.db.QueryRowContext(ctx, `SELECT COUNT(*) FROM referrals WHERE inviter_id=? AND status IN ('reg_rewarded','paid_rewarded')`, userID).Scan(®) + s.db.QueryRowContext(ctx, `SELECT COUNT(*) FROM referrals WHERE inviter_id=? AND status='paid_rewarded'`, userID).Scan(&paid) + // 邀请人:每 reg_rewarded +3、每 paid_rewarded 再 +7;被邀请人自身得的天数不计入其「邀请战绩」。 + earnedDays = tg + reg*3 + paid*7 + return +} + +func isDup(err error) bool { + if err == nil { + return false + } + m := strings.ToLower(err.Error()) + return strings.Contains(m, "unique") || strings.Contains(m, "duplicate") || strings.Contains(m, "constraint") +} diff --git a/server/internal/reward/store_sqlite_test.go b/server/internal/reward/store_sqlite_test.go new file mode 100644 index 0000000..1e34e2f --- /dev/null +++ b/server/internal/reward/store_sqlite_test.go @@ -0,0 +1,76 @@ +package reward + +import ( + "context" + "database/sql" + "testing" + "time" + + "github.com/wangjia/pangolin/server/internal/config" + "github.com/wangjia/pangolin/server/internal/store" +) + +func openDB(t *testing.T) *sql.DB { + db, err := store.Open(&config.Config{Driver: "sqlite", DSN: ":memory:"}) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = db.Close() }) + if err := store.MigrateUp(db, "sqlite"); err != nil { + t.Fatal(err) + } + if err := store.ApplyCodesLibMigrations(context.Background(), db, "sqlite"); err != nil { + t.Fatal(err) + } + return db +} + +func seedU(t *testing.T, db *sql.DB, id int64, uuid string) { + _, err := db.Exec(`INSERT INTO users (id,uuid,email,pw_hash,dp_uuid,status,created_at) + VALUES (?,?,?, 'x','dp-'||?, 'active', ?)`, id, uuid, uuid+"@x", uuid, time.Now().UTC()) + if err != nil { + t.Fatal(err) + } +} + +func TestEnsureAndResolveInviteCode(t *testing.T) { + db := openDB(t) + seedU(t, db, 1, "u1") + st := NewStore(db) + code, err := st.EnsureInviteCode(context.Background(), 1, func() string { return "ABCD2345" }) + if err != nil || code != "ABCD2345" { + t.Fatalf("ensure: %q %v", code, err) + } + again, _ := st.EnsureInviteCode(context.Background(), 1, func() string { return "ZZZZ9999" }) + if again != "ABCD2345" { + t.Fatalf("second ensure changed code: %q", again) + } + inviter, ok, _ := st.ResolveInviteCode(context.Background(), "ABCD2345") + if !ok || inviter != 1 { + t.Fatalf("resolve: %d %v", inviter, ok) + } +} + +func TestInsertClaimUniqueGuards(t *testing.T) { + db := openDB(t) + seedU(t, db, 1, "u1") + seedU(t, db, 2, "u2") + st := NewStore(db) + tx, _ := db.Begin() + if err := st.InsertClaimTx(context.Background(), tx, 1, "telegram_join", "tg-100", 3, time.Now().UTC()); err != nil { + t.Fatal(err) + } + tx.Commit() + // 同 user 再领 → ErrClaimExists + tx2, _ := db.Begin() + if err := st.InsertClaimTx(context.Background(), tx2, 1, "telegram_join", "tg-999", 3, time.Now().UTC()); err != ErrClaimExists { + t.Fatalf("same user reclaim err = %v, want ErrClaimExists", err) + } + tx2.Rollback() + // 同 telegram_id 换 user → ErrClaimExists + tx3, _ := db.Begin() + if err := st.InsertClaimTx(context.Background(), tx3, 2, "telegram_join", "tg-100", 3, time.Now().UTC()); err != ErrClaimExists { + t.Fatalf("same tgid reclaim err = %v, want ErrClaimExists", err) + } + tx3.Rollback() +}