* 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>
354 lines
14 KiB
Go
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)
|
|
}
|