Files
agent/backend/internal/seed/demo_postgres_test.go

184 lines
7.2 KiB
Go

package seed
import (
"context"
"fmt"
"os"
"sync"
"sync/atomic"
"testing"
"time"
"github.com/stretchr/testify/require"
"golang.org/x/crypto/bcrypt"
"gorm.io/driver/postgres"
"gorm.io/gorm"
"gorm.io/gorm/logger"
"senlinai-agent/backend/internal/models"
)
func TestPostgresDemoSeedSerializesConcurrentRunsForSameOwner(t *testing.T) {
dsn := os.Getenv("DATABASE_URL")
if dsn == "" {
t.Skip("DATABASE_URL is not configured; skipping PostgreSQL seed concurrency test")
}
database, err := gorm.Open(postgres.Open(dsn), &gorm.Config{
TranslateError: true,
Logger: logger.Default.LogMode(logger.Silent),
})
require.NoError(t, err)
require.NoError(t, models.AutoMigrate(database))
suffix := fmt.Sprint(time.Now().UnixNano())
email := "postgres-seed-" + suffix + "@example.com"
passwordHash, err := bcrypt.GenerateFromPassword([]byte("original-password"), bcrypt.MinCost)
require.NoError(t, err)
user := models.SenlinAgentUser{Email: email, DisplayName: "保留用户名称", PasswordHash: string(passwordHash), Role: "user"}
require.NoError(t, database.Create(&user).Error)
projects := []models.SenlinAgentProject{
{OwnerID: user.ID, Name: "项目 A1", Identifier: "A1"},
{OwnerID: user.ID, Name: "项目 A2", Identifier: "A2"},
{OwnerID: user.ID, Name: "数据中台", Identifier: "DATA"},
{OwnerID: user.ID, Name: "运营自动化", Identifier: "AUTO"},
}
for index := range projects {
require.NoError(t, database.Create(&projects[index]).Error)
}
uiTag := models.SenlinAgentTag{ProjectID: projects[0].ID, Name: "UI"}
require.NoError(t, database.Create(&uiTag).Error)
require.NoError(t, database.Create(&models.SenlinAgentTask{
ProjectID: projects[0].ID, CreatedBy: user.ID, TagID: &uiTag.ID,
Title: "智能报表导出功能", Status: "open",
}).Error)
var ownerLockObserved atomic.Bool
var tagMisses atomic.Int32
releaseTagMisses := make(chan struct{})
var releaseOnce sync.Once
callbackName := "test:coordinate_concurrent_seed_" + suffix
require.NoError(t, database.Callback().Query().After("gorm:query").Register(callbackName, func(tx *gorm.DB) {
if tx.Statement.Schema == nil {
return
}
if tx.Statement.Schema.Table == (models.SenlinAgentUser{}).TableName() {
if _, ok := tx.Statement.Clauses["FOR"]; ok {
ownerLockObserved.Store(true)
}
return
}
if ownerLockObserved.Load() || tx.Statement.Schema.Table != (models.SenlinAgentTag{}).TableName() || tx.RowsAffected != 0 {
return
}
if tagMisses.Add(1) >= 2 {
releaseOnce.Do(func() { close(releaseTagMisses) })
}
select {
case <-releaseTagMisses:
case <-time.After(5 * time.Second):
tx.AddError(fmt.Errorf("timed out coordinating concurrent tag misses"))
}
}))
t.Cleanup(func() {
database.Callback().Query().Remove(callbackName)
cleanupPostgresDemoSeed(database, user.ID)
})
start := make(chan struct{})
results := make(chan error, 2)
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
var wait sync.WaitGroup
for range 2 {
wait.Add(1)
go func() {
defer wait.Done()
<-start
_, err := Demo(database.WithContext(ctx), DemoOptions{Email: email, DisplayName: "不应覆盖", Password: "new-password"})
results <- err
}()
}
close(start)
done := make(chan struct{})
go func() {
wait.Wait()
close(done)
}()
select {
case <-done:
case <-ctx.Done():
<-done
t.Fatal("timed out waiting for concurrent Demo calls")
}
close(results)
for err := range results {
require.NoError(t, err)
}
require.True(t, ownerLockObserved.Load(), "seed must lock the owner user row before creating child objects")
var storedUser models.SenlinAgentUser
require.NoError(t, database.Where("email = ?", email).First(&storedUser).Error)
require.Equal(t, "保留用户名称", storedUser.DisplayName)
require.NoError(t, bcrypt.CompareHashAndPassword([]byte(storedUser.PasswordHash), []byte("original-password")))
require.Error(t, bcrypt.CompareHashAndPassword([]byte(storedUser.PasswordHash), []byte("new-password")))
projectIDs := []uint{projects[0].ID, projects[1].ID, projects[2].ID, projects[3].ID}
assertPostgresDemoCount(t, database, &models.SenlinAgentProject{}, "owner_id = ?", []any{user.ID}, 4)
assertPostgresDemoCount(t, database, &models.SenlinAgentTag{}, "project_id = ?", []any{projects[0].ID}, 4)
assertPostgresDemoCount(t, database, &models.SenlinAgentTask{}, "project_id IN ?", []any{projectIDs}, 6)
assertPostgresDemoCount(t, database, &models.SenlinAgentInboxItem{}, "project_id IN ?", []any{projectIDs}, 7)
assertPostgresDemoCount(t, database, &models.SenlinAgentAISession{}, "project_id = ?", []any{projects[0].ID}, 4)
assertPostgresDemoCount(t, database, &models.SenlinAgentNote{}, "project_id = ?", []any{projects[0].ID}, 2)
assertPostgresDemoCount(t, database, &models.SenlinAgentSource{}, "project_id = ?", []any{projects[0].ID}, 2)
assertPostgresDemoCount(t, database, &models.SenlinAgentCronPlan{}, "project_id = ?", []any{projects[0].ID}, 3)
assertPostgresDemoCount(t, database, &models.SenlinAgentProjectChannel{}, "project_id = ?", []any{projects[0].ID}, 2)
}
func assertPostgresDemoCount(t *testing.T, database *gorm.DB, model any, where string, args []any, expected int64) {
t.Helper()
var count int64
require.NoError(t, database.Model(model).Where(where, args...).Count(&count).Error)
require.Equal(t, expected, count)
}
func cleanupPostgresDemoSeed(database *gorm.DB, userID uint) {
var projects []models.SenlinAgentProject
if database.Where("owner_id = ?", userID).Find(&projects).Error != nil {
return
}
projectIDs := make([]uint, 0, len(projects))
for _, project := range projects {
projectIDs = append(projectIDs, project.ID)
}
if len(projectIDs) != 0 {
var tasks []models.SenlinAgentTask
database.Where("project_id IN ?", projectIDs).Find(&tasks)
taskIDs := make([]uint, 0, len(tasks))
for _, task := range tasks {
taskIDs = append(taskIDs, task.ID)
}
if len(taskIDs) != 0 {
database.Where("task_id IN ?", taskIDs).Delete(&models.SenlinAgentTaskShare{})
}
var inboxItems []models.SenlinAgentInboxItem
database.Where("project_id IN ?", projectIDs).Find(&inboxItems)
inboxIDs := make([]uint, 0, len(inboxItems))
for _, item := range inboxItems {
inboxIDs = append(inboxIDs, item.ID)
}
if len(inboxIDs) != 0 {
database.Where("inbox_item_id IN ?", inboxIDs).Delete(&models.SenlinAgentInboxSuggestion{})
}
database.Where("project_id IN ?", projectIDs).Delete(&models.SenlinAgentProjectEvent{})
database.Where("project_id IN ?", projectIDs).Delete(&models.SenlinAgentTask{})
database.Where("project_id IN ?", projectIDs).Delete(&models.SenlinAgentNote{})
database.Where("project_id IN ?", projectIDs).Delete(&models.SenlinAgentSource{})
database.Where("project_id IN ?", projectIDs).Delete(&models.SenlinAgentAISession{})
database.Where("project_id IN ?", projectIDs).Delete(&models.SenlinAgentTag{})
database.Where("project_id IN ?", projectIDs).Delete(&models.SenlinAgentProjectChannel{})
database.Where("project_id IN ?", projectIDs).Delete(&models.SenlinAgentCronPlan{})
database.Where("project_id IN ?", projectIDs).Delete(&models.SenlinAgentInboxItem{})
database.Where("id IN ?", projectIDs).Delete(&models.SenlinAgentProject{})
}
database.Delete(&models.SenlinAgentUser{}, userID)
}