211 lines
7.9 KiB
Go
211 lines
7.9 KiB
Go
package seed
|
|
|
|
import (
|
|
"fmt"
|
|
"sync/atomic"
|
|
"testing"
|
|
"time"
|
|
|
|
"github.com/glebarez/sqlite"
|
|
"github.com/stretchr/testify/require"
|
|
"golang.org/x/crypto/bcrypt"
|
|
"gorm.io/gorm"
|
|
"gorm.io/gorm/logger"
|
|
"senlinai-agent/backend/internal/logic/auth"
|
|
"senlinai-agent/backend/internal/logic/projects"
|
|
"senlinai-agent/backend/internal/models"
|
|
)
|
|
|
|
func TestDemoSeedCreatesFrontendWorkspaceData(t *testing.T) {
|
|
database := newTestDB(t)
|
|
|
|
result, err := Demo(database, DemoOptions{
|
|
Email: "demo@senlin.ai",
|
|
DisplayName: "演示用户",
|
|
Password: "password123",
|
|
})
|
|
|
|
require.NoError(t, err)
|
|
require.Equal(t, "demo@senlin.ai", result.User.Email)
|
|
require.Len(t, result.Projects, 4)
|
|
|
|
token, err := auth.NewService("test-secret").Login("demo@senlin.ai", "password123")
|
|
require.NoError(t, err)
|
|
require.NotEmpty(t, token)
|
|
|
|
workspace, err := projects.NewService().Workspace(result.User.ID, result.Projects[0].Identity)
|
|
require.NoError(t, err)
|
|
require.Equal(t, "项目 A1", workspace.Project.Name)
|
|
require.GreaterOrEqual(t, len(workspace.Inbox), 4)
|
|
require.GreaterOrEqual(t, len(workspace.Tasks), 3)
|
|
require.GreaterOrEqual(t, len(workspace.AISessions), 4)
|
|
require.GreaterOrEqual(t, len(workspace.NotesSources), 3)
|
|
require.GreaterOrEqual(t, len(workspace.CronPlans), 3)
|
|
require.Len(t, workspace.Channels, 8)
|
|
}
|
|
|
|
func TestDemoSeedIsIdempotent(t *testing.T) {
|
|
database := newTestDB(t)
|
|
|
|
first, err := Demo(database, DemoOptions{Email: "demo@senlin.ai", DisplayName: "演示用户", Password: "password123"})
|
|
require.NoError(t, err)
|
|
second, err := Demo(database, DemoOptions{Email: "demo@senlin.ai", DisplayName: "演示用户", Password: "password123"})
|
|
require.NoError(t, err)
|
|
|
|
require.Equal(t, first.User.ID, second.User.ID)
|
|
require.Len(t, second.Projects, 4)
|
|
var userCount int64
|
|
require.NoError(t, database.Model(&models.SenlinAgentUser{}).Where("email = ?", "demo@senlin.ai").Count(&userCount).Error)
|
|
require.Equal(t, int64(1), userCount)
|
|
|
|
var projectCount int64
|
|
require.NoError(t, database.Model(&models.SenlinAgentProject{}).Where("owner_id = ?", first.User.ID).Count(&projectCount).Error)
|
|
require.Equal(t, int64(4), projectCount)
|
|
var distinctIdentifiers int64
|
|
require.NoError(t, database.Model(&models.SenlinAgentProject{}).
|
|
Where("owner_id = ?", first.User.ID).
|
|
Distinct("identifier").
|
|
Count(&distinctIdentifiers).Error)
|
|
require.Equal(t, projectCount, distinctIdentifiers)
|
|
|
|
var tags []models.SenlinAgentTag
|
|
require.NoError(t, database.Where("project_id = ?", first.Projects[0].ID).Find(&tags).Error)
|
|
require.Len(t, tags, 4)
|
|
for _, tag := range tags {
|
|
require.Equal(t, first.Projects[0].ID, tag.ProjectID)
|
|
require.NotEmpty(t, tag.Identity)
|
|
}
|
|
}
|
|
|
|
func TestDemoSeedUsesIdentifierAsStableProjectKey(t *testing.T) {
|
|
database := newTestDB(t)
|
|
|
|
first, err := Demo(database, DemoOptions{Email: "demo@senlin.ai", DisplayName: "演示用户", Password: "password123"})
|
|
require.NoError(t, err)
|
|
original := first.Projects[0]
|
|
require.NoError(t, database.Model(&original).Update("name", "本地重命名项目").Error)
|
|
|
|
second, err := Demo(database, DemoOptions{Email: "demo@senlin.ai", DisplayName: "演示用户", Password: "password123"})
|
|
|
|
require.NoError(t, err)
|
|
require.Equal(t, original.ID, second.Projects[0].ID)
|
|
require.Equal(t, "A1", second.Projects[0].Identifier)
|
|
require.Equal(t, "本地重命名项目", second.Projects[0].Name)
|
|
var projectCount int64
|
|
require.NoError(t, database.Model(&models.SenlinAgentProject{}).
|
|
Where("owner_id = ? AND identifier = ?", first.User.ID, "A1").
|
|
Count(&projectCount).Error)
|
|
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) {
|
|
database := newTestDB(t)
|
|
|
|
first, err := Demo(database, DemoOptions{Email: "demo@senlin.ai", DisplayName: "演示用户", Password: "password123"})
|
|
require.NoError(t, err)
|
|
otherTag := models.SenlinAgentTag{ProjectID: first.Projects[1].ID, Name: "UI"}
|
|
require.NoError(t, database.Create(&otherTag).Error)
|
|
|
|
_, err = Demo(database, DemoOptions{Email: "demo@senlin.ai", DisplayName: "演示用户", Password: "password123"})
|
|
|
|
require.NoError(t, err)
|
|
for _, projectID := range []uint{first.Projects[0].ID, first.Projects[1].ID} {
|
|
var count int64
|
|
require.NoError(t, database.Model(&models.SenlinAgentTag{}).
|
|
Where("project_id = ? AND name = ?", projectID, "UI").
|
|
Count(&count).Error)
|
|
require.Equal(t, int64(1), count)
|
|
}
|
|
|
|
var taggedTasks []models.SenlinAgentTask
|
|
require.NoError(t, database.Where("project_id = ? AND tag_id IS NOT NULL", first.Projects[0].ID).Find(&taggedTasks).Error)
|
|
require.NotEmpty(t, taggedTasks)
|
|
for _, task := range taggedTasks {
|
|
var tag models.SenlinAgentTag
|
|
require.NoError(t, database.First(&tag, *task.TagID).Error)
|
|
require.Equal(t, task.ProjectID, tag.ProjectID)
|
|
}
|
|
}
|
|
|
|
func newTestDB(t *testing.T) *gorm.DB {
|
|
t.Helper()
|
|
database, err := gorm.Open(sqlite.Open(fmt.Sprintf("file:%s?mode=memory&cache=shared", t.Name())), &gorm.Config{
|
|
Logger: logger.Default.LogMode(logger.Silent),
|
|
})
|
|
require.NoError(t, err)
|
|
require.NoError(t, models.AutoMigrate(database))
|
|
models.DBService = database
|
|
return database
|
|
}
|