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) }