topfans/backend/services/activityService/repository/activity_repository_test.go
2026-06-24 18:10:12 +08:00

162 lines
4.9 KiB
Go

package repository
import (
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/topfans/backend/pkg/database"
"github.com/topfans/backend/pkg/models"
"gorm.io/gorm"
)
// setupTestDB 设置测试数据库(与 assetService 共用同一实例)
func setupTestDB(t *testing.T) *gorm.DB {
config := database.Config{
Host: "localhost",
Port: 5432,
User: "haihuizhu",
Password: "admin",
DBName: "top-fans",
SSLMode: "disable",
TimeZone: "Asia/Shanghai",
}
if err := database.Init(config); err != nil {
t.Skipf("Skipping test: failed to connect to test database: %v", err)
}
db := database.GetDB()
if err := db.AutoMigrate(&models.Activity{}, &models.ActivityUserStats{}, &models.ActivityContribution{}); err != nil {
t.Logf("Warning: failed to migrate activity tables: %v", err)
}
cleanupActivityTestDB(t, db)
return db
}
// cleanupActivityTestDB 清理测试数据
func cleanupActivityTestDB(_ *testing.T, db *gorm.DB) {
db.Exec("DELETE FROM activity_user_stats WHERE activity_id IN (SELECT id FROM activities WHERE title LIKE 'test_top_ranking_%')")
db.Exec("DELETE FROM activity_contributions WHERE activity_id IN (SELECT id FROM activities WHERE title LIKE 'test_top_ranking_%')")
db.Exec("DELETE FROM activities WHERE title LIKE 'test_top_ranking_%'")
}
// createTestActivity 创建测试活动
func createTestActivity(t *testing.T, db *gorm.DB, title string, starID int64) *models.Activity {
now := time.Now().Unix()
act := &models.Activity{
Title: title,
StarID: starID,
Status: "active",
StartTime: now - 3600,
EndTime: now + 3600,
CreatedAt: now,
UpdatedAt: now,
}
if err := db.Create(act).Error; err != nil {
t.Fatalf("Failed to create test activity: %v", err)
}
return act
}
// createTestStats 创建用户活动统计
func createTestStats(t *testing.T, db *gorm.DB, activityID, userID, starID, totalContribution int64) *models.ActivityUserStats {
now := time.Now().Unix()
stats := &models.ActivityUserStats{
ActivityID: activityID,
UserID: userID,
StarID: starID,
TotalContribution: totalContribution,
TotalCrystalSpent: 0,
TotalItems: 0,
LastContributeAt: now,
CreatedAt: now,
UpdatedAt: now,
}
if err := db.Create(stats).Error; err != nil {
t.Fatalf("Failed to create test stats: %v", err)
}
return stats
}
// TestGetTop3_Empty 测试活动没有任何 stats 时返回空切片
func TestGetTop3_Empty(t *testing.T) {
db := setupTestDB(t)
defer cleanupActivityTestDB(t, db)
repo := NewActivityRepository()
act := createTestActivity(t, db, "test_top_ranking_empty", 9999)
stats, err := repo.GetTop3(act.ID, 0)
assert.NoError(t, err)
assert.Empty(t, stats)
}
// TestGetTop3_LessThan3 测试只有 1 行时返回 1 行
func TestGetTop3_LessThan3(t *testing.T) {
db := setupTestDB(t)
defer cleanupActivityTestDB(t, db)
repo := NewActivityRepository()
act := createTestActivity(t, db, "test_top_ranking_one", 9999)
createTestStats(t, db, act.ID, 1001, 9999, 500)
stats, err := repo.GetTop3(act.ID, 0)
assert.NoError(t, err)
assert.Len(t, stats, 1)
assert.Equal(t, int64(1001), stats[0].UserID)
assert.Equal(t, int64(500), stats[0].TotalContribution)
}
// TestGetTop3_FullWithStar 测试 3 行且带 star_id 过滤
func TestGetTop3_FullWithStar(t *testing.T) {
db := setupTestDB(t)
defer cleanupActivityTestDB(t, db)
repo := NewActivityRepository()
act := createTestActivity(t, db, "test_top_ranking_full", 7777)
// 制造 5 条:3 条 star=7777, 2 条 star=8888;应只返回 star=7777 的前 3
createTestStats(t, db, act.ID, 1001, 7777, 900)
createTestStats(t, db, act.ID, 1002, 7777, 800)
createTestStats(t, db, act.ID, 1003, 7777, 700)
createTestStats(t, db, act.ID, 2001, 8888, 9999)
createTestStats(t, db, act.ID, 2002, 8888, 9998)
stats, err := repo.GetTop3(act.ID, 7777)
assert.NoError(t, err)
assert.Len(t, stats, 3)
assert.Equal(t, int64(1001), stats[0].UserID)
assert.Equal(t, int64(1002), stats[1].UserID)
assert.Equal(t, int64(1003), stats[2].UserID)
}
// TestGetUserStatsForRanking_NotFound 测试未找到时返回 (nil, nil)
func TestGetUserStatsForRanking_NotFound(t *testing.T) {
db := setupTestDB(t)
defer cleanupActivityTestDB(t, db)
repo := NewActivityRepository()
act := createTestActivity(t, db, "test_top_ranking_notfound", 6666)
stats, err := repo.GetUserStatsForRanking(act.ID, 99999, 0)
assert.NoError(t, err)
assert.Nil(t, stats)
}
// TestGetUserStatsForRanking_Found 测试找到时返回正确 stats
func TestGetUserStatsForRanking_Found(t *testing.T) {
db := setupTestDB(t)
defer cleanupActivityTestDB(t, db)
repo := NewActivityRepository()
act := createTestActivity(t, db, "test_top_ranking_found", 6666)
createTestStats(t, db, act.ID, 1001, 6666, 1500)
stats, err := repo.GetUserStatsForRanking(act.ID, 1001, 6666)
assert.NoError(t, err)
assert.NotNil(t, stats)
assert.Equal(t, int64(1500), stats.TotalContribution)
assert.Equal(t, int64(1001), stats.UserID)
}