package codes_test import ( "context" "database/sql" "errors" "testing" "time" libcodes "github.com/wangjia/codes" "github.com/wangjia/pangolin/server/internal/codes" ) // TestCreateBatchViaLib verifies CreateBatch mints codes through the shared // codes library (Mint): codes land in the lib's new `codes` table (no // plaintext persisted), the batch's entitlement is a duration payload, and // unknown plans are still rejected up front (old behaviour, service.go:336-339). func TestCreateBatchViaLib(t *testing.T) { ctx := context.Background() db := openMigratedSQLite(t) store := codes.NewStore(db) svc := codes.NewService(store, nil, 5, time.Hour) // rdb=nil: cmd/codegen's real usage res, err := svc.CreateBatch(ctx, codes.BatchRequest{ PlanCode: codes.PlanPro, DurationDays: 30, Count: 5, Channel: codes.ChannelManual, Note: "t", CreatedBy: "admin:1", }) if err != nil { t.Fatalf("CreateBatch: %v", err) } if len(res.Codes) != 5 || res.BatchID == 0 { t.Fatalf("res = %+v", res) } // 码落在库的新表里,明文不落库,权益为 duration payload。 var cnt int if err := db.QueryRow(`SELECT COUNT(*) FROM codes WHERE batch_id=?`, res.BatchID).Scan(&cnt); err != nil || cnt != 5 { t.Fatalf("new codes rows = %d (err=%v), want 5", cnt, err) } var kind, payload string if err := db.QueryRow(`SELECT entitlement_kind, entitlement_payload FROM codes_batches WHERE id=?`, res.BatchID).Scan(&kind, &payload); err != nil { t.Fatalf("batch: %v", err) } if kind != "duration" || payload != `{"plan":"pro","days":30}` { t.Fatalf("kind=%q payload=%s", kind, payload) } for _, plain := range res.Codes { var n int if err := db.QueryRow(`SELECT COUNT(*) FROM codes WHERE code_hash=?`, codes.Hash(plain)).Scan(&n); err != nil || n != 1 { t.Fatalf("hash lookup for %q: n=%d err=%v", plain, n, err) } } // 未知 plan 仍报错(旧行为 service.go:336-339)。 if _, err := svc.CreateBatch(ctx, codes.BatchRequest{ PlanCode: "nope", DurationDays: 30, Count: 1, Channel: codes.ChannelManual, CreatedBy: "x", }); err == nil { t.Fatal("unknown plan should error") } } // -------------------------------------------------------------------------- // Redeem — via GuardedRedeem, grant callback in the same tx as the flip. // -------------------------------------------------------------------------- func newRedeemFixture(t *testing.T) (*sql.DB, *codes.Service, string) { t.Helper() db := openMigratedSQLite(t) seedUser(t, db, 1) seedUser(t, db, 2) store := codes.NewStore(db) svc := codes.NewService(store, nil, 5, time.Hour) res, err := svc.CreateBatch(context.Background(), codes.BatchRequest{ PlanCode: codes.PlanPro, DurationDays: 30, Count: 1, Channel: codes.ChannelManual, CreatedBy: "t", }) if err != nil { t.Fatalf("mint: %v", err) } return db, svc, res.Codes[0] } func TestRedeemCreatesSubscription(t *testing.T) { ctx := context.Background() db, svc, plain := newRedeemFixture(t) r, apiErr := svc.Redeem(ctx, codes.RedeemRequest{UserID: 1, Code: plain}) if apiErr != nil { t.Fatalf("redeem: %v", apiErr) } if r.Idempotent || r.PlanCode != codes.PlanPro || r.DurationDays != 30 || r.SubscriptionID == 0 { t.Fatalf("result = %+v", r) } wantMin := time.Now().UTC().AddDate(0, 0, 29) if r.ExpiresAt.Before(wantMin) { t.Errorf("expires %v, want ≥ ~30d out", r.ExpiresAt) } // 订阅真落库(source='code'),码翻转 + redeemed_by 是 "user:1"。 var src string var exp time.Time if err := db.QueryRow(`SELECT source, expires_at FROM subscriptions WHERE id=?`, r.SubscriptionID).Scan(&src, &exp); err != nil || src != "code" { t.Fatalf("sub: src=%q err=%v", src, err) } var status, rby string if err := db.QueryRow(`SELECT status, redeemed_by FROM codes WHERE code_hash=?`, codes.Hash(mustCanonical(t, plain))).Scan(&status, &rby); err != nil { t.Fatalf("code: %v", err) } if status != "redeemed" || rby != "user:1" { t.Errorf("status=%q redeemed_by=%q", status, rby) } // 旧式审计行仍写 audit_log(admin 审计页数据源)。 var audits int if err := db.QueryRow(`SELECT COUNT(*) FROM audit_log WHERE action='redeem' AND actor='user:1'`).Scan(&audits); err != nil || audits != 1 { t.Errorf("audit_log redeem rows = %d (err=%v), want 1", audits, err) } } // TestWebhookMintsIntoLibTables verifies Store.MintOne (webhook 铸码核心): // lands in the shared lib's new tables, entitlement is a duration payload, // and a duplicate code_hash leaves no orphan batch row (the old // two-statement webhook implementation leaked one batch per dupe delivery). func TestWebhookMintsIntoLibTables(t *testing.T) { ctx := context.Background() db := openMigratedSQLite(t) store := codes.NewStore(db) plain, err := codes.GenerateCode() if err != nil { t.Fatal(err) } canonical := mustCanonical(t, plain) ent, err := libcodesEnt("pro", 30) if err != nil { t.Fatal(err) } if err := store.MintOne(ctx, codes.Hash(canonical), ent, codes.ChannelStore, "webhook", "n1"); err != nil { t.Fatalf("MintOne: %v", err) } var batches, ccount int if err := db.QueryRow(`SELECT COUNT(*) FROM codes_batches`).Scan(&batches); err != nil || batches != 1 { t.Fatalf("batches = %d (err=%v)", batches, err) } if err := db.QueryRow(`SELECT COUNT(*) FROM codes WHERE code_hash=?`, codes.Hash(canonical)).Scan(&ccount); err != nil || ccount != 1 { t.Fatalf("codes = %d (err=%v)", ccount, err) } // 同 hash 再灌:ErrDuplicate 且 batch 数不变(无孤儿——旧实现的 I2 缺陷在此修复)。 err = store.MintOne(ctx, codes.Hash(canonical), ent, codes.ChannelStore, "webhook", "n2") if !errors.Is(err, codes.ErrDuplicate) { t.Fatalf("err = %v, want ErrDuplicate", err) } if err := db.QueryRow(`SELECT COUNT(*) FROM codes_batches`).Scan(&batches); err != nil || batches != 1 { t.Fatalf("batches after dup = %d, want 1 (no orphan)", batches) } } // libcodesEnt: 测试侧小工具,避免测试文件 import 库包名与宿主包名混淆。 func libcodesEnt(plan string, days int) (libcodes.Entitlement, error) { return libcodes.NewDurationEntitlement(plan, days) } func mustCanonical(t *testing.T, code string) string { t.Helper() c, err := codes.Canonicalize(code) if err != nil { t.Fatalf("canonicalize: %v", err) } return c } func TestRedeemIdempotentSameUser(t *testing.T) { ctx := context.Background() _, svc, plain := newRedeemFixture(t) if _, apiErr := svc.Redeem(ctx, codes.RedeemRequest{UserID: 1, Code: plain}); apiErr != nil { t.Fatalf("first: %v", apiErr) } r2, apiErr := svc.Redeem(ctx, codes.RedeemRequest{UserID: 1, Code: plain}) if apiErr != nil { t.Fatalf("second: %v", apiErr) } if !r2.Idempotent || r2.PlanCode != codes.PlanPro || r2.DurationDays != 30 { t.Fatalf("r2 = %+v", r2) } if !r2.ExpiresAt.IsZero() { t.Errorf("idempotent replay must omit ExpiresAt (old contract), got %v", r2.ExpiresAt) } } func TestRedeemOtherUserRejected(t *testing.T) { ctx := context.Background() _, svc, plain := newRedeemFixture(t) if _, apiErr := svc.Redeem(ctx, codes.RedeemRequest{UserID: 1, Code: plain}); apiErr != nil { t.Fatalf("first: %v", apiErr) } _, apiErr := svc.Redeem(ctx, codes.RedeemRequest{UserID: 2, Code: plain}) if apiErr == nil || apiErr.Code != "CODE_REDEEMED" { t.Fatalf("apiErr = %+v, want CODE_REDEEMED", apiErr) } } func TestRedeemSamePlanExtends(t *testing.T) { ctx := context.Background() db, svc, plain := newRedeemFixture(t) // 预置同 plan 活跃订阅,到期在未来 10 天 → 兑换后应为 +10+30 天(max(expires,now)+days)。 var proID int64 if err := db.QueryRow(`SELECT id FROM plans WHERE code='pro'`).Scan(&proID); err != nil { t.Fatalf("plan: %v", err) } base := time.Now().UTC().AddDate(0, 0, 10).Truncate(time.Second) var subID int64 res, err := db.Exec(`INSERT INTO subscriptions (user_id, plan_id, expires_at, source, created_at) VALUES (1, ?, ?, 'code', ?)`, proID, base, time.Now().UTC()) if err != nil { t.Fatalf("seed sub: %v", err) } subID, _ = res.LastInsertId() r, apiErr := svc.Redeem(ctx, codes.RedeemRequest{UserID: 1, Code: plain}) if apiErr != nil { t.Fatalf("redeem: %v", apiErr) } if r.SubscriptionID != subID { t.Fatalf("extended sub id = %d, want %d (extend, not create)", r.SubscriptionID, subID) } want := base.AddDate(0, 0, 30) if !r.ExpiresAt.Equal(want) { t.Errorf("expires = %v, want %v", r.ExpiresAt, want) } var cnt int if err := db.QueryRow(`SELECT COUNT(*) FROM subscriptions WHERE user_id=1`).Scan(&cnt); err != nil || cnt != 1 { t.Errorf("subs = %d, want 1 (extended in place)", cnt) } } func TestRedeemCrossPlanCreatesNew(t *testing.T) { ctx := context.Background() db, svc, plain := newRedeemFixture(t) // 码是 pro var teamID int64 if err := db.QueryRow(`SELECT id FROM plans WHERE code='team'`).Scan(&teamID); err != nil { t.Fatalf("plan: %v", err) } if _, err := db.Exec(`INSERT INTO subscriptions (user_id, plan_id, expires_at, source, created_at) VALUES (1, ?, ?, 'code', ?)`, teamID, time.Now().UTC().AddDate(0, 0, 90), time.Now().UTC()); err != nil { t.Fatalf("seed team sub: %v", err) } r, apiErr := svc.Redeem(ctx, codes.RedeemRequest{UserID: 1, Code: plain}) if apiErr != nil { t.Fatalf("redeem: %v", apiErr) } var cnt int if err := db.QueryRow(`SELECT COUNT(*) FROM subscriptions WHERE user_id=1`).Scan(&cnt); err != nil || cnt != 2 { t.Fatalf("subs = %d, want 2 (new pro sub beside team)", cnt) } wantMin := time.Now().UTC().AddDate(0, 0, 29) if r.ExpiresAt.Before(wantMin) { t.Errorf("new pro sub expires %v, want ~30d (NOT stacked on team's 90d)", r.ExpiresAt) } } func TestRedeemInvalidFormat(t *testing.T) { ctx := context.Background() _, svc, _ := newRedeemFixture(t) _, apiErr := svc.Redeem(ctx, codes.RedeemRequest{UserID: 1, Code: "not-a-code"}) if apiErr == nil || apiErr.Code != "INVALID_CODE" { t.Fatalf("apiErr = %+v, want INVALID_CODE", apiErr) } } func TestRedeemNotFound(t *testing.T) { ctx := context.Background() _, svc, _ := newRedeemFixture(t) // 结构合法但不存在的码:生成一个不入库的。 plain, err := codes.GenerateCode() if err != nil { t.Fatal(err) } _, apiErr := svc.Redeem(ctx, codes.RedeemRequest{UserID: 1, Code: plain}) if apiErr == nil || apiErr.Code != "CODE_NOT_FOUND" { t.Fatalf("apiErr = %+v, want CODE_NOT_FOUND", apiErr) } }