H-60: harden Captain migration rollback and concurrency (#10)
Co-authored-by: Rogee <rogee@ipao.vip>
This commit is contained in:
@@ -3,13 +3,18 @@ package service
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/url"
|
||||
"os"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gochat/gochat/internal/model"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/driver/sqlite"
|
||||
"gorm.io/driver/postgres"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/logger"
|
||||
)
|
||||
@@ -42,31 +47,74 @@ func TestWidgetConversationRollsBackCaptainBindingOnCreateFailure(t *testing.T)
|
||||
}
|
||||
|
||||
func TestEnsureCaptainAgentBotBindingConcurrentCallsStayUnique(t *testing.T) {
|
||||
db, err := gorm.Open(sqlite.Open("file:"+t.Name()+"?mode=memory&cache=shared"), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)})
|
||||
if os.Getenv("GOCHAT_TEST_DB") == "sqlite" {
|
||||
t.Skip("requires PostgreSQL conflict handling")
|
||||
}
|
||||
dsn := os.Getenv("GOCHAT_TEST_DB_URL")
|
||||
if dsn == "" {
|
||||
dsn = "host=localhost port=5432 user=postgres password=postgres dbname=gochat_test sslmode=disable"
|
||||
}
|
||||
admin, err := gorm.Open(postgres.Open(dsn), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)})
|
||||
require.NoError(t, err)
|
||||
schema := fmt.Sprintf("captain_binding_%d", time.Now().UnixNano())
|
||||
require.NoError(t, admin.Exec("CREATE SCHEMA "+schema).Error)
|
||||
t.Cleanup(func() {
|
||||
_ = admin.Exec("DROP SCHEMA " + schema + " CASCADE").Error
|
||||
if sqlDB, dbErr := admin.DB(); dbErr == nil {
|
||||
_ = sqlDB.Close()
|
||||
}
|
||||
})
|
||||
if strings.HasPrefix(dsn, "postgres://") || strings.HasPrefix(dsn, "postgresql://") {
|
||||
dsnURL, parseErr := url.Parse(dsn)
|
||||
require.NoError(t, parseErr)
|
||||
query := dsnURL.Query()
|
||||
query.Set("search_path", schema)
|
||||
dsnURL.RawQuery = query.Encode()
|
||||
dsn = dsnURL.String()
|
||||
} else {
|
||||
dsn += " search_path=" + schema
|
||||
}
|
||||
db, err := gorm.Open(postgres.Open(dsn), &gorm.Config{
|
||||
Logger: logger.Default.LogMode(logger.Silent),
|
||||
DisableForeignKeyConstraintWhenMigrating: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, db.AutoMigrate(&model.CaptainAssistant{}, &model.AgentBot{}, &model.AgentBotInbox{}))
|
||||
sqlDB, err := db.DB()
|
||||
require.NoError(t, err)
|
||||
sqlDB.SetMaxOpenConns(1)
|
||||
const calls = 8
|
||||
sqlDB.SetMaxOpenConns(calls)
|
||||
t.Cleanup(func() { _ = sqlDB.Close() })
|
||||
assistant := &model.CaptainAssistant{AccountID: 1, Name: "Concurrent", Status: model.AssistantStatusActive}
|
||||
require.NoError(t, db.Create(assistant).Error)
|
||||
|
||||
const calls = 8
|
||||
errs := make(chan error, calls)
|
||||
ids := make(chan uint, calls)
|
||||
ready := make(chan struct{}, calls)
|
||||
start := make(chan struct{})
|
||||
var wg sync.WaitGroup
|
||||
for range calls {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
bot, err := ensureCaptainAgentBotBinding(context.Background(), db, assistant, 42)
|
||||
var bot *model.AgentBot
|
||||
err := db.Transaction(func(tx *gorm.DB) error {
|
||||
ready <- struct{}{}
|
||||
<-start
|
||||
var bindErr error
|
||||
bot, bindErr = ensureCaptainAgentBotBinding(context.Background(), tx, assistant, 42)
|
||||
return bindErr
|
||||
})
|
||||
errs <- err
|
||||
if bot != nil {
|
||||
ids <- bot.ID
|
||||
}
|
||||
}()
|
||||
}
|
||||
for range calls {
|
||||
<-ready
|
||||
}
|
||||
close(start)
|
||||
wg.Wait()
|
||||
close(errs)
|
||||
close(ids)
|
||||
|
||||
Reference in New Issue
Block a user