Files
gochat/backend/internal/service/captain_skill_service_test.go
T
Rogeeandrogee aaad337ef3 H-292: prove Captain Skill updates reach CAS (#54)
* H-292: prove Captain Skill updates reach CAS

* H-292: prove handler conflict comes from CAS

* H-292: reuse profile test password digest

---------

Co-authored-by: Rogee <rogee@ipao.vip>
2026-08-19 17:55:19 +08:00

354 lines
14 KiB
Go

package service
import (
"context"
"errors"
"fmt"
"os"
"path/filepath"
"strings"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/driver/postgres"
"gorm.io/driver/sqlite"
"gorm.io/gorm"
"gorm.io/gorm/logger"
"github.com/gochat/gochat/internal/model"
"github.com/gochat/gochat/internal/repository"
)
func newCaptainSkillTestService(t *testing.T) (*CaptainSkillService, *gorm.DB, *model.CaptainAssistant, *model.CaptainAssistant) {
t.Helper()
db, err := gorm.Open(sqlite.Open("file:"+filepath.Join(t.TempDir(), "captain-skill.db")+"?_busy_timeout=5000&_journal_mode=WAL"), &gorm.Config{
Logger: logger.Default.LogMode(logger.Silent),
})
require.NoError(t, err)
sqlDB, err := db.DB()
require.NoError(t, err)
sqlDB.SetMaxOpenConns(2)
t.Cleanup(func() { require.NoError(t, sqlDB.Close()) })
require.NoError(t, db.AutoMigrate(
&model.CaptainAssistant{},
&model.CaptainSkill{},
&model.CaptainSkillReference{},
&model.CaptainAssistantSkill{},
))
assistants := []*model.CaptainAssistant{
{AccountID: 1, Name: "Account one", Status: model.AssistantStatusActive},
{AccountID: 2, Name: "Account two", Status: model.AssistantStatusActive},
}
for _, assistant := range assistants {
require.NoError(t, db.Create(assistant).Error)
}
return NewCaptainSkillService(repository.NewCaptainSkillRepo(db)), db, assistants[0], assistants[1]
}
func validCaptainSkillRequest() *CaptainSkillRequest {
return &CaptainSkillRequest{
Name: "refund-policy",
Description: "Refund rules",
InstructionsMD: "Use the relevant reference.",
Status: model.CaptainSkillStatusActive,
ExpectedVersion: 1,
References: []CaptainSkillReferenceRequest{
{ReferenceKey: "regional-exceptions", ContentMD: "Region A differs."},
{ReferenceKey: "standard-policy", ContentMD: "Standard terms."},
},
}
}
func TestCaptainSkillServiceManagementFlow(t *testing.T) {
svc, db, assistant, otherTenantAssistant := newCaptainSkillTestService(t)
ctx := context.Background()
skill, err := svc.Create(ctx, assistant.AccountID, validCaptainSkillRequest())
require.NoError(t, err)
assert.Equal(t, uint(1), skill.Version)
require.Len(t, skill.References, 2)
assert.Equal(t, 0, skill.References[0].Position)
firstReferenceID := skill.References[0].ID
_, err = svc.Get(ctx, otherTenantAssistant.AccountID, skill.ID)
assert.ErrorIs(t, err, gorm.ErrRecordNotFound)
assert.ErrorIs(t, svc.Bind(ctx, assistant.AccountID, otherTenantAssistant.ID, skill.ID), gorm.ErrRecordNotFound)
assert.ErrorIs(t, svc.Bind(ctx, otherTenantAssistant.AccountID, otherTenantAssistant.ID, skill.ID), gorm.ErrRecordNotFound)
require.NoError(t, svc.Bind(ctx, assistant.AccountID, assistant.ID, skill.ID))
require.NoError(t, svc.Bind(ctx, assistant.AccountID, assistant.ID, skill.ID))
list, err := svc.List(ctx, assistant.AccountID, &assistant.ID)
require.NoError(t, err)
require.Len(t, list, 1)
assert.True(t, list[0].Bound)
assert.Equal(t, int64(1), list[0].BoundAssistantCount)
assert.Equal(t, int64(2), list[0].ReferenceCount)
request := validCaptainSkillRequest()
request.Description = "Updated rules"
request.References = request.References[:1]
updated, err := svc.Update(ctx, assistant.AccountID, skill.ID, request)
require.NoError(t, err)
assert.Equal(t, uint(2), updated.Version)
require.Len(t, updated.References, 1)
assert.Equal(t, firstReferenceID, updated.References[0].ID, "reference_key reconciliation preserves IDs")
request.ExpectedVersion = updated.Version
unchanged, err := svc.Update(ctx, assistant.AccountID, skill.ID, request)
require.NoError(t, err)
assert.Equal(t, uint(2), unchanged.Version)
request.Status = model.CaptainSkillStatusArchived
archived, err := svc.Update(ctx, assistant.AccountID, skill.ID, request)
require.NoError(t, err)
assert.Equal(t, uint(3), archived.Version)
list, err = svc.List(ctx, assistant.AccountID, &assistant.ID)
require.NoError(t, err)
assert.True(t, list[0].Bound, "archiving preserves the binding")
request.ExpectedVersion = archived.Version
request.Status = model.CaptainSkillStatusActive
active, err := svc.Update(ctx, assistant.AccountID, skill.ID, request)
require.NoError(t, err)
assert.Equal(t, uint(4), active.Version)
assert.ErrorIs(t, svc.Delete(ctx, assistant.AccountID, skill.ID), ErrCaptainSkillConflict)
require.NoError(t, svc.Unbind(ctx, assistant.AccountID, assistant.ID, skill.ID))
require.NoError(t, svc.Unbind(ctx, assistant.AccountID, assistant.ID, skill.ID))
request.ExpectedVersion = active.Version
request.Status = model.CaptainSkillStatusDraft
draft, err := svc.Update(ctx, assistant.AccountID, skill.ID, request)
require.NoError(t, err)
assert.Equal(t, uint(5), draft.Version)
require.NoError(t, svc.Delete(ctx, assistant.AccountID, skill.ID))
assert.ErrorIs(t, db.First(&model.CaptainSkill{}, skill.ID).Error, gorm.ErrRecordNotFound)
}
func TestCaptainSkillServiceConcurrentUpdateConflict(t *testing.T) {
svc, db, assistant, _ := newCaptainSkillTestService(t)
skill, err := svc.Create(context.Background(), assistant.AccountID, validCaptainSkillRequest())
require.NoError(t, err)
t.Run("content and references", func(t *testing.T) {
assertConcurrentCaptainSkillUpdate(t, db, svc, skill, false)
})
t.Run("status only", func(t *testing.T) {
draftRequest := validCaptainSkillRequest()
draftRequest.Name = "status-only"
draftRequest.Status = model.CaptainSkillStatusDraft
draft, err := svc.Create(context.Background(), assistant.AccountID, draftRequest)
require.NoError(t, err)
assertConcurrentCaptainSkillUpdate(t, db, svc, draft, true)
})
t.Run("repository CAS", func(t *testing.T) {
request := validCaptainSkillRequest()
request.Name = "repository-cas"
current, err := svc.Create(context.Background(), assistant.AccountID, request)
require.NoError(t, err)
next := *current
next.Version++
require.NoError(t, svc.repo.Update(context.Background(), &next, current.Version, false, false))
assert.ErrorIs(t, svc.repo.Update(context.Background(), &next, current.Version, false, false), repository.ErrCaptainSkillVersionConflict)
})
}
func TestCaptainSkillServicePostgresSmoke(t *testing.T) {
if os.Getenv("GOCHAT_TEST_DB") == "sqlite" {
t.Skip("PostgreSQL service test")
}
dsn := os.Getenv("GOCHAT_TEST_DB_URL")
if dsn == "" {
dsn = "host=localhost port=5432 user=postgres password=postgres dbname=gochat_test sslmode=disable"
}
db, err := gorm.Open(postgres.Open(dsn), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)})
require.NoError(t, err)
adminDB := db
schema := fmt.Sprintf("captain_skill_service_%d", time.Now().UnixNano())
require.NoError(t, adminDB.Exec("CREATE SCHEMA "+schema).Error)
t.Cleanup(func() { _ = adminDB.Exec("DROP SCHEMA " + schema + " CASCADE").Error })
if strings.Contains(dsn, "://") {
separator := "?"
if strings.Contains(dsn, "?") {
separator = "&"
}
dsn += separator + "search_path=" + schema
} else {
dsn += " search_path=" + schema
}
db, err = gorm.Open(postgres.Open(dsn), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)})
require.NoError(t, err)
require.NoError(t, db.AutoMigrate(&model.CaptainAssistant{}, &model.CaptainSkill{}, &model.CaptainSkillReference{}, &model.CaptainAssistantSkill{}))
assistant := &model.CaptainAssistant{AccountID: 1, Name: "Postgres", Status: model.AssistantStatusActive}
require.NoError(t, db.Create(assistant).Error)
svc := NewCaptainSkillService(repository.NewCaptainSkillRepo(db))
skill, err := svc.Create(context.Background(), 1, validCaptainSkillRequest())
require.NoError(t, err)
update := validCaptainSkillRequest()
update.References[0], update.References[1] = update.References[1], update.References[0]
update.References[0].ContentMD = "Updated and reordered."
skill, err = svc.Update(context.Background(), 1, skill.ID, update)
require.NoError(t, err)
assert.Equal(t, uint(2), skill.Version)
assert.Equal(t, "standard-policy", skill.References[0].ReferenceKey)
assertConcurrentCaptainSkillUpdate(t, db, svc, skill, false)
require.NoError(t, svc.Bind(context.Background(), 1, assistant.ID, skill.ID))
items, err := svc.List(context.Background(), 1, &assistant.ID)
require.NoError(t, err)
require.Len(t, items, 1)
assert.True(t, items[0].Bound)
draftRequest := validCaptainSkillRequest()
draftRequest.Name = "status-only"
draftRequest.Status = model.CaptainSkillStatusDraft
draft, err := svc.Create(context.Background(), 1, draftRequest)
require.NoError(t, err)
assertConcurrentCaptainSkillUpdate(t, db, svc, draft, true)
}
func assertConcurrentCaptainSkillUpdate(t *testing.T, db *gorm.DB, svc *CaptainSkillService, skill *model.CaptainSkill, statusOnly bool) {
t.Helper()
ready := make(chan struct{}, 2)
release := make(chan struct{})
callbackName := "test:captain_skill_cas_barrier:" + t.Name()
require.NoError(t, db.Callback().Update().Before("gorm:update").Register(callbackName, func(tx *gorm.DB) {
if tx.Statement.Table == "captain_skills" {
ready <- struct{}{}
<-release
}
}))
defer func() { require.NoError(t, db.Callback().Update().Remove(callbackName)) }()
type updateResult struct {
skill *model.CaptainSkill
err error
}
results := make(chan updateResult, 2)
start := make(chan struct{})
for i := 0; i < 2; i++ {
req := validCaptainSkillRequest()
req.Name = skill.Name
req.ExpectedVersion = skill.Version
if statusOnly {
req.Status = []model.CaptainSkillStatus{model.CaptainSkillStatusActive, model.CaptainSkillStatusArchived}[i]
} else {
req.Description = fmt.Sprintf("Concurrent update %d", i)
req.References[0].ContentMD = fmt.Sprintf("Concurrent reference %d", i)
}
go func() {
<-start
updated, err := svc.Update(context.Background(), skill.AccountID, skill.ID, req)
results <- updateResult{updated, err}
}()
}
close(start)
for range 2 {
select {
case <-ready:
case <-time.After(5 * time.Second):
close(release)
t.Fatal("concurrent updates did not both reach the repository CAS")
}
}
close(release)
var winner *model.CaptainSkill
conflicts := 0
for i := 0; i < 2; i++ {
result := <-results
if errors.Is(result.err, ErrCaptainSkillConflict) {
assert.ErrorContains(t, result.err, repository.ErrCaptainSkillVersionConflict.Error())
conflicts++
continue
}
require.NoError(t, result.err)
require.Nil(t, winner, "only one concurrent update may succeed")
winner = result.skill
}
require.Equal(t, 1, conflicts)
require.NotNil(t, winner)
persisted, err := svc.Get(context.Background(), skill.AccountID, skill.ID)
require.NoError(t, err)
assert.Equal(t, skill.Version+1, persisted.Version)
if statusOnly {
assert.Equal(t, winner.Status, persisted.Status)
} else {
assert.Equal(t, winner.Description, persisted.Description)
assert.Equal(t, winner.References[0].ContentMD, persisted.References[0].ContentMD)
}
}
func TestCaptainSkillServiceValidationAndRollback(t *testing.T) {
svc, db, assistant, _ := newCaptainSkillTestService(t)
ctx := context.Background()
tests := []struct {
name string
mutate func(*CaptainSkillRequest)
}{
{"invalid status", func(r *CaptainSkillRequest) { r.Status = "invalid" }},
{"unsafe reference key", func(r *CaptainSkillRequest) { r.References[0].ReferenceKey = "../secret" }},
{"duplicate reference key", func(r *CaptainSkillRequest) { r.References[1].ReferenceKey = r.References[0].ReferenceKey }},
{"instructions too large", func(r *CaptainSkillRequest) {
r.InstructionsMD = strings.Repeat("x", CaptainSkillMaxInstructionsBytes+1)
}},
{"reference too large", func(r *CaptainSkillRequest) {
r.References[0].ContentMD = strings.Repeat("x", CaptainSkillMaxReferenceBytes+1)
}},
{"too many references", func(r *CaptainSkillRequest) {
r.References = make([]CaptainSkillReferenceRequest, CaptainSkillMaxReferences+1)
for i := range r.References {
r.References[i] = CaptainSkillReferenceRequest{ReferenceKey: "key-" + string(rune('a'+i)), ContentMD: "x"}
}
}},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
req := validCaptainSkillRequest()
tt.mutate(req)
_, err := svc.Create(ctx, assistant.AccountID, req)
assert.ErrorIs(t, err, ErrCaptainSkillValidation)
})
}
draft := validCaptainSkillRequest()
draft.Status = model.CaptainSkillStatusDraft
skill, err := svc.Create(ctx, assistant.AccountID, draft)
require.NoError(t, err)
assert.ErrorIs(t, svc.Bind(ctx, assistant.AccountID, assistant.ID, skill.ID), ErrCaptainSkillConflict)
for i := 0; i < 50; i++ {
limited := &model.CaptainSkill{AccountID: assistant.AccountID, Name: fmt.Sprintf("limit-%d", i), Description: "x", InstructionsMD: "x", Status: model.CaptainSkillStatusActive, Version: 1}
require.NoError(t, db.Create(limited).Error)
require.NoError(t, db.Create(&model.CaptainAssistantSkill{AccountID: assistant.AccountID, AssistantID: assistant.ID, SkillID: limited.ID}).Error)
}
overLimit := validCaptainSkillRequest()
overLimit.Name = "over-limit"
overLimitSkill, err := svc.Create(ctx, assistant.AccountID, overLimit)
require.NoError(t, err)
assert.ErrorIs(t, svc.Bind(ctx, assistant.AccountID, assistant.ID, overLimitSkill.ID), ErrCaptainSkillConflict)
archived := validCaptainSkillRequest()
archived.Name = "archived-bound"
archived.Status = model.CaptainSkillStatusArchived
archivedSkill, err := svc.Create(ctx, assistant.AccountID, archived)
require.NoError(t, err)
require.NoError(t, db.Create(&model.CaptainAssistantSkill{AccountID: assistant.AccountID, AssistantID: assistant.ID, SkillID: archivedSkill.ID}).Error)
archived.Status = model.CaptainSkillStatusActive
_, err = svc.Update(ctx, assistant.AccountID, archivedSkill.ID, archived)
assert.ErrorIs(t, err, ErrCaptainSkillConflict)
callbackName := "test:captain_skill_reference_failure"
require.NoError(t, db.Callback().Create().Before("gorm:create").Register(callbackName, func(tx *gorm.DB) {
if tx.Statement.Table == "captain_skill_references" {
require.ErrorContains(t, tx.AddError(errors.New("forced reference failure")), "forced reference failure")
}
}))
t.Cleanup(func() { _ = db.Callback().Create().Remove(callbackName) })
update := validCaptainSkillRequest()
update.Description = "must roll back"
update.References[0].ContentMD = "must also roll back"
_, err = svc.Update(ctx, assistant.AccountID, skill.ID, update)
require.ErrorContains(t, err, "forced reference failure")
persisted, err := svc.Get(ctx, assistant.AccountID, skill.ID)
require.NoError(t, err)
assert.Equal(t, uint(1), persisted.Version)
assert.NotEqual(t, "must roll back", persisted.Description)
require.Len(t, persisted.References, 2)
}