fix: serialize concurrent demo seeding
This commit is contained in:
@@ -1,14 +1,18 @@
|
|||||||
package seed
|
package seed
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"errors"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"golang.org/x/crypto/bcrypt"
|
"golang.org/x/crypto/bcrypt"
|
||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
|
"gorm.io/gorm/clause"
|
||||||
"senlinai-agent/backend/internal/models"
|
"senlinai-agent/backend/internal/models"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
const demoTransactionAttempts = 3
|
||||||
|
|
||||||
type DemoOptions struct {
|
type DemoOptions struct {
|
||||||
Email string
|
Email string
|
||||||
DisplayName string
|
DisplayName string
|
||||||
@@ -35,8 +39,11 @@ func Demo(database *gorm.DB, options DemoOptions) (DemoResult, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
var result DemoResult
|
var result DemoResult
|
||||||
err := database.Transaction(func(tx *gorm.DB) error {
|
var err error
|
||||||
user, err := upsertUser(tx, email, displayName, password)
|
for attempt := 0; attempt < demoTransactionAttempts; attempt++ {
|
||||||
|
result = DemoResult{}
|
||||||
|
err = database.Transaction(func(tx *gorm.DB) error {
|
||||||
|
user, err := upsertAndLockUser(tx, email, displayName, password)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -49,29 +56,51 @@ func Demo(database *gorm.DB, options DemoOptions) (DemoResult, error) {
|
|||||||
result.Projects = projects
|
result.Projects = projects
|
||||||
return nil
|
return nil
|
||||||
})
|
})
|
||||||
|
if err == nil || !isRetryableDemoTransactionError(err) || attempt == demoTransactionAttempts-1 {
|
||||||
|
return result, err
|
||||||
|
}
|
||||||
|
// 重试发生在事务完全回滚之后,避免持锁等待自身的新事务。
|
||||||
|
time.Sleep(time.Duration(attempt+1) * 10 * time.Millisecond)
|
||||||
|
}
|
||||||
return result, err
|
return result, err
|
||||||
}
|
}
|
||||||
|
|
||||||
func upsertUser(tx *gorm.DB, email string, displayName string, password string) (models.SenlinAgentUser, error) {
|
func upsertAndLockUser(tx *gorm.DB, email string, displayName string, password string) (models.SenlinAgentUser, error) {
|
||||||
var user models.SenlinAgentUser
|
|
||||||
err := tx.Where("email = ?", email).First(&user).Error
|
|
||||||
if err == nil {
|
|
||||||
return user, nil
|
|
||||||
}
|
|
||||||
if err != gorm.ErrRecordNotFound {
|
|
||||||
return user, err
|
|
||||||
}
|
|
||||||
hash, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost)
|
hash, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return user, err
|
return models.SenlinAgentUser{}, err
|
||||||
}
|
}
|
||||||
user = models.SenlinAgentUser{
|
candidate := models.SenlinAgentUser{
|
||||||
Email: email,
|
Email: email,
|
||||||
DisplayName: displayName,
|
DisplayName: displayName,
|
||||||
PasswordHash: string(hash),
|
PasswordHash: string(hash),
|
||||||
Role: "user",
|
Role: "user",
|
||||||
}
|
}
|
||||||
return user, tx.Create(&user).Error
|
if err := tx.Clauses(clause.OnConflict{
|
||||||
|
Columns: []clause.Column{{Name: "email"}},
|
||||||
|
DoNothing: true,
|
||||||
|
}).Create(&candidate).Error; err != nil {
|
||||||
|
return models.SenlinAgentUser{}, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// PostgreSQL 会先等待并发的同 email INSERT 提交,再在 READ COMMITTED 的下一条语句中看见该用户。
|
||||||
|
// 锁定 owner 后,项目、标签和子对象的 find-or-create 在不同进程间也会按同一用户串行执行。
|
||||||
|
var user models.SenlinAgentUser
|
||||||
|
err = tx.Clauses(clause.Locking{Strength: "UPDATE"}).Where("email = ?", email).First(&user).Error
|
||||||
|
return user, err
|
||||||
|
}
|
||||||
|
|
||||||
|
func isRetryableDemoTransactionError(err error) bool {
|
||||||
|
var sqlState interface{ SQLState() string }
|
||||||
|
if !errors.As(err, &sqlState) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
switch sqlState.SQLState() {
|
||||||
|
case "40001", "40P01":
|
||||||
|
return true
|
||||||
|
default:
|
||||||
|
return false
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func seedProjects(tx *gorm.DB, ownerID uint) ([]models.SenlinAgentProject, error) {
|
func seedProjects(tx *gorm.DB, ownerID uint) ([]models.SenlinAgentProject, error) {
|
||||||
|
|||||||
183
backend/internal/seed/demo_postgres_test.go
Normal file
183
backend/internal/seed/demo_postgres_test.go
Normal file
@@ -0,0 +1,183 @@
|
|||||||
|
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)
|
||||||
|
}
|
||||||
@@ -2,10 +2,13 @@ package seed
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"sync/atomic"
|
||||||
"testing"
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
"github.com/glebarez/sqlite"
|
"github.com/glebarez/sqlite"
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
|
"golang.org/x/crypto/bcrypt"
|
||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
"gorm.io/gorm/logger"
|
"gorm.io/gorm/logger"
|
||||||
"senlinai-agent/backend/internal/logic/auth"
|
"senlinai-agent/backend/internal/logic/auth"
|
||||||
@@ -87,6 +90,7 @@ func TestDemoSeedUsesIdentifierAsStableProjectKey(t *testing.T) {
|
|||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.Equal(t, original.ID, second.Projects[0].ID)
|
require.Equal(t, original.ID, second.Projects[0].ID)
|
||||||
require.Equal(t, "A1", second.Projects[0].Identifier)
|
require.Equal(t, "A1", second.Projects[0].Identifier)
|
||||||
|
require.Equal(t, "本地重命名项目", second.Projects[0].Name)
|
||||||
var projectCount int64
|
var projectCount int64
|
||||||
require.NoError(t, database.Model(&models.SenlinAgentProject{}).
|
require.NoError(t, database.Model(&models.SenlinAgentProject{}).
|
||||||
Where("owner_id = ? AND identifier = ?", first.User.ID, "A1").
|
Where("owner_id = ? AND identifier = ?", first.User.ID, "A1").
|
||||||
@@ -94,6 +98,77 @@ func TestDemoSeedUsesIdentifierAsStableProjectKey(t *testing.T) {
|
|||||||
require.Equal(t, int64(1), projectCount)
|
require.Equal(t, int64(1), projectCount)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestDemoSeedAcceptsAConcurrentWinnerWithoutOverwritingUser(t *testing.T) {
|
||||||
|
database := newTestDB(t)
|
||||||
|
const email = "seed-race@senlin.ai"
|
||||||
|
winnerHash, err := bcrypt.GenerateFromPassword([]byte("winner-password"), bcrypt.MinCost)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
var injected atomic.Bool
|
||||||
|
callbackName := "test:inject_concurrent_seed_user"
|
||||||
|
require.NoError(t, database.Callback().Create().Before("gorm:create").Register(callbackName, func(tx *gorm.DB) {
|
||||||
|
if tx.Statement.Schema == nil || tx.Statement.Schema.Table != (models.SenlinAgentUser{}).TableName() {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if !injected.CompareAndSwap(false, true) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
now := time.Now().UTC()
|
||||||
|
tx.AddError(tx.Exec(
|
||||||
|
"INSERT INTO senlin_agent_users (identity, email, display_name, password_hash, role, created_at, updated_at) VALUES (?, ?, ?, ?, ?, ?, ?)",
|
||||||
|
"00000000-0000-7000-8000-000000000001", email, "并发胜出用户", string(winnerHash), "user", now, now,
|
||||||
|
).Error)
|
||||||
|
}))
|
||||||
|
t.Cleanup(func() { database.Callback().Create().Remove(callbackName) })
|
||||||
|
|
||||||
|
result, err := Demo(database, DemoOptions{Email: email, DisplayName: "不应覆盖", Password: "loser-password"})
|
||||||
|
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.True(t, injected.Load())
|
||||||
|
require.Equal(t, "并发胜出用户", result.User.DisplayName)
|
||||||
|
require.NoError(t, bcrypt.CompareHashAndPassword([]byte(result.User.PasswordHash), []byte("winner-password")))
|
||||||
|
require.Error(t, bcrypt.CompareHashAndPassword([]byte(result.User.PasswordHash), []byte("loser-password")))
|
||||||
|
var userCount int64
|
||||||
|
require.NoError(t, database.Model(&models.SenlinAgentUser{}).Where("email = ?", email).Count(&userCount).Error)
|
||||||
|
require.Equal(t, int64(1), userCount)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDemoSeedRetriesSerializationFailureOutsideTheTransaction(t *testing.T) {
|
||||||
|
database := newTestDB(t)
|
||||||
|
var failures atomic.Int32
|
||||||
|
callbackName := "test:fail_first_seed_user_create"
|
||||||
|
require.NoError(t, database.Callback().Create().Before("gorm:create").Register(callbackName, func(tx *gorm.DB) {
|
||||||
|
if tx.Statement.Schema == nil || tx.Statement.Schema.Table != (models.SenlinAgentUser{}).TableName() {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if failures.Add(1) == 1 {
|
||||||
|
tx.AddError(testSQLStateError{state: "40001"})
|
||||||
|
}
|
||||||
|
}))
|
||||||
|
t.Cleanup(func() { database.Callback().Create().Remove(callbackName) })
|
||||||
|
|
||||||
|
result, err := Demo(database, DemoOptions{Email: "retry@senlin.ai", DisplayName: "重试用户", Password: "password123"})
|
||||||
|
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Equal(t, int32(2), failures.Load())
|
||||||
|
require.Equal(t, "retry@senlin.ai", result.User.Email)
|
||||||
|
var userCount int64
|
||||||
|
require.NoError(t, database.Model(&models.SenlinAgentUser{}).Where("email = ?", result.User.Email).Count(&userCount).Error)
|
||||||
|
require.Equal(t, int64(1), userCount)
|
||||||
|
}
|
||||||
|
|
||||||
|
type testSQLStateError struct {
|
||||||
|
state string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (err testSQLStateError) Error() string {
|
||||||
|
return "test PostgreSQL transaction failure " + err.state
|
||||||
|
}
|
||||||
|
|
||||||
|
func (err testSQLStateError) SQLState() string {
|
||||||
|
return err.state
|
||||||
|
}
|
||||||
|
|
||||||
func TestDemoSeedKeepsSameNamedTagsScopedToTheirProjects(t *testing.T) {
|
func TestDemoSeedKeepsSameNamedTagsScopedToTheirProjects(t *testing.T) {
|
||||||
database := newTestDB(t)
|
database := newTestDB(t)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user