Files
pangolin/server/internal/notices/store_sqlite_test.go
T

97 lines
3.0 KiB
Go

package notices
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 {
t.Helper()
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)
}
_ = store.ApplyCodesLibMigrations(context.Background(), db, "sqlite")
return db
}
func seedU(t *testing.T, db *sql.DB, id int64, uuid string) {
t.Helper()
if _, 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()); err != nil {
t.Fatal(err)
}
}
func TestListForUser_MergeFilterUnread(t *testing.T) {
db := openDB(t)
seedU(t, db, 1, "u1")
seedU(t, db, 2, "u2")
st := NewStore(db)
now := time.Now().UTC()
// 广播1条 + 我的定向1条 + 他人定向1条 + 已撤回1条 + 已过期1条
if _, err := st.InsertBroadcast(context.Background(), "news", "b1", "b1", "", "", "", now.Add(-3*time.Hour), nil); err != nil {
t.Fatal(err)
}
tx, _ := db.Begin()
if err := st.InsertNoticeTx(context.Background(), tx, 1, "reward", "mine", "mine", "", "", "", now.Add(-2*time.Hour)); err != nil {
t.Fatal(err)
}
if err := st.InsertNoticeTx(context.Background(), tx, 2, "reward", "other", "other", "", "", "", now.Add(-1*time.Hour)); err != nil {
t.Fatal(err)
}
_ = tx.Commit()
rid, _ := st.InsertBroadcast(context.Background(), "promo", "revoked", "revoked", "", "", "", now, nil)
_ = st.Revoke(context.Background(), rid, now)
exp := now.Add(-time.Minute)
_, _ = st.InsertBroadcast(context.Background(), "promo", "expired", "expired", "", "", "", now.Add(-4*time.Hour), &exp)
items, unread, err := st.ListForUser(context.Background(), 1, now, 50)
if err != nil {
t.Fatal(err)
}
if len(items) != 2 {
t.Fatalf("items=%d want 2(广播+我的;他人/撤回/过期均排除): %+v", len(items), items)
}
if items[0].TitleZH != "mine" || items[1].TitleZH != "b1" {
t.Fatalf("排序应 published_at 倒序: %+v", items)
}
if unread != 2 {
t.Fatalf("水位为空→全未读, unread=%d", unread)
}
// 置水位到 -2.5h:广播(-3h)已读、定向(-2h)未读
if err := st.MarkRead(context.Background(), 1, now.Add(-150*time.Minute)); err != nil {
t.Fatal(err)
}
items, unread, _ = st.ListForUser(context.Background(), 1, now, 50)
if unread != 1 || !items[0].Unread || items[1].Unread {
t.Fatalf("水位判定错: unread=%d items=%+v", unread, items)
}
}
func TestInsertNoticeTx_RollsBackWithTx(t *testing.T) {
db := openDB(t)
seedU(t, db, 1, "u1")
st := NewStore(db)
tx, _ := db.Begin()
if err := st.InsertNoticeTx(context.Background(), tx, 1, "reward", "t", "t", "", "", "", time.Now().UTC()); err != nil {
t.Fatal(err)
}
_ = tx.Rollback()
var n int
_ = db.QueryRow(`SELECT COUNT(*) FROM notices`).Scan(&n)
if n != 0 {
t.Fatalf("回滚后应零孤儿通知, got %d", n)
}
}