Files
gochat/internal/repository/captain_scenario_repo_test.go
T
2026-06-04 15:44:48 +08:00

217 lines
6.7 KiB
Go

package repository
import (
"context"
"encoding/json"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/gochat/gochat/internal/model"
)
// createTestScenario is a helper that builds a CaptainScenario with typical fields.
func createTestScenario(accountID uint, assistantID uint, title string) *model.CaptainScenario {
tools, _ := json.Marshal([]string{"handoff", "search"})
return &model.CaptainScenario{
AccountID: accountID,
AssistantID: assistantID,
Title: title,
Description: "A test scenario for " + title,
Instruction: "Please handle " + title + " inquiries",
Enabled: true,
Tools: tools,
}
}
// --- 1. Create ---
func TestCaptainScenarioRepo_Create(t *testing.T) {
db := setupTestDB(t, &model.CaptainScenario{})
repo := NewCaptainScenarioRepo(db)
tools, _ := json.Marshal([]string{"handoff", "search"})
scenario := &model.CaptainScenario{
AccountID: 1,
AssistantID: 10,
Title: "TestScenario",
Description: "A test scenario",
Instruction: "Handle test inquiries",
Enabled: true,
Tools: tools,
}
err := repo.Create(context.Background(), scenario)
require.NoError(t, err)
assert.NotZero(t, scenario.ID, "ID should be set after Create")
assert.Equal(t, "TestScenario", scenario.Title)
assert.Equal(t, uint(1), scenario.AccountID)
assert.Equal(t, uint(10), scenario.AssistantID)
assert.True(t, scenario.Enabled)
// Verify Tools JSON was persisted
var foundTools []string
require.NoError(t, json.Unmarshal(scenario.Tools, &foundTools))
assert.Equal(t, []string{"handoff", "search"}, foundTools)
}
// --- 2. GetByID ---
func TestCaptainScenarioRepo_GetByID(t *testing.T) {
db := setupTestDB(t, &model.CaptainScenario{})
repo := NewCaptainScenarioRepo(db)
scenario := createTestScenario(1, 10, "GetByIDScenario")
err := repo.Create(context.Background(), scenario)
require.NoError(t, err)
found, err := repo.GetByID(context.Background(), scenario.ID)
require.NoError(t, err)
assert.Equal(t, scenario.ID, found.ID)
assert.Equal(t, "GetByIDScenario", found.Title)
assert.Equal(t, uint(1), found.AccountID)
assert.Equal(t, uint(10), found.AssistantID)
assert.True(t, found.Enabled)
// Verify Tools JSON round-tripped
var foundTools []string
require.NoError(t, json.Unmarshal(found.Tools, &foundTools))
assert.Equal(t, []string{"handoff", "search"}, foundTools)
}
// --- 3. GetByID Not Found ---
func TestCaptainScenarioRepo_GetByID_NotFound(t *testing.T) {
db := setupTestDB(t, &model.CaptainScenario{})
repo := NewCaptainScenarioRepo(db)
found, err := repo.GetByID(context.Background(), 9999)
assert.Error(t, err)
assert.Nil(t, found)
}
// --- 4. Update ---
func TestCaptainScenarioRepo_Update(t *testing.T) {
db := setupTestDB(t, &model.CaptainScenario{})
repo := NewCaptainScenarioRepo(db)
scenario := createTestScenario(1, 10, "BeforeUpdate")
err := repo.Create(context.Background(), scenario)
require.NoError(t, err)
// Modify fields
scenario.Title = "AfterUpdate"
scenario.Description = "Updated description"
scenario.Instruction = "Updated instruction"
scenario.Enabled = false
newTools, _ := json.Marshal([]string{"handoff", "weather"})
scenario.Tools = newTools
err = repo.Update(context.Background(), scenario)
require.NoError(t, err)
found, err := repo.GetByID(context.Background(), scenario.ID)
require.NoError(t, err)
assert.Equal(t, "AfterUpdate", found.Title)
assert.Equal(t, "Updated description", found.Description)
assert.Equal(t, "Updated instruction", found.Instruction)
assert.False(t, found.Enabled)
// Verify Tools JSON was updated
var foundTools []string
require.NoError(t, json.Unmarshal(found.Tools, &foundTools))
assert.Equal(t, []string{"handoff", "weather"}, foundTools)
}
// --- 5. Delete ---
func TestCaptainScenarioRepo_Delete(t *testing.T) {
db := setupTestDB(t, &model.CaptainScenario{})
repo := NewCaptainScenarioRepo(db)
scenario := createTestScenario(1, 10, "DeleteMe")
err := repo.Create(context.Background(), scenario)
require.NoError(t, err)
err = repo.Delete(context.Background(), scenario.ID)
require.NoError(t, err)
// After deletion, GetByID should return error
found, err := repo.GetByID(context.Background(), scenario.ID)
assert.Error(t, err)
assert.Nil(t, found)
}
// --- 6. ListByAssistant ---
func TestCaptainScenarioRepo_ListByAssistant(t *testing.T) {
db := setupTestDB(t, &model.CaptainScenario{})
repo := NewCaptainScenarioRepo(db)
// Create 3 scenarios under assistant 10
for i := 0; i < 3; i++ {
s := createTestScenario(1, 10, "ListScenario"+string(rune('A'+i)))
require.NoError(t, repo.Create(context.Background(), s))
}
// Create 1 scenario under assistant 20
s2 := createTestScenario(1, 20, "OtherAssistantScenario")
require.NoError(t, repo.Create(context.Background(), s2))
scenarios, count, err := repo.ListByAssistant(context.Background(), 10, 0, 10)
require.NoError(t, err)
assert.Equal(t, int64(3), count)
assert.Len(t, scenarios, 3)
// Verify each scenario belongs to assistant 10
for _, s := range scenarios {
assert.Equal(t, uint(10), s.AssistantID)
}
// Pagination: offset=1, limit=2 should return 2 items (total still 3)
scenarios2, count2, err := repo.ListByAssistant(context.Background(), 10, 1, 2)
require.NoError(t, err)
assert.Equal(t, int64(3), count2)
assert.Len(t, scenarios2, 2)
}
// --- 7. ListByAssistant Empty ---
func TestCaptainScenarioRepo_ListByAssistant_Empty(t *testing.T) {
db := setupTestDB(t, &model.CaptainScenario{})
repo := NewCaptainScenarioRepo(db)
scenarios, count, err := repo.ListByAssistant(context.Background(), 9999, 0, 10)
require.NoError(t, err)
assert.Equal(t, int64(0), count)
assert.Len(t, scenarios, 0)
}
// --- 8. FindEnabled ---
func TestCaptainScenarioRepo_FindEnabled(t *testing.T) {
skipIfSQLite(t) // WHERE enabled = true relies on PG boolean semantics
db := setupTestDB(t, &model.CaptainScenario{})
repo := NewCaptainScenarioRepo(db)
// Create 2 enabled scenarios
s1 := createTestScenario(1, 10, "EnabledA")
s1.Enabled = true
require.NoError(t, repo.Create(context.Background(), s1))
s2 := createTestScenario(1, 10, "EnabledB")
s2.Enabled = true
require.NoError(t, repo.Create(context.Background(), s2))
// Create 1 disabled scenario
s3 := createTestScenario(1, 10, "DisabledC")
s3.Enabled = false
require.NoError(t, repo.Create(context.Background(), s3))
// Create 1 enabled scenario under different assistant
s4 := createTestScenario(1, 20, "OtherEnabled")
s4.Enabled = true
require.NoError(t, repo.Create(context.Background(), s4))
enabled, err := repo.FindEnabled(context.Background(), 10)
require.NoError(t, err)
assert.Len(t, enabled, 2)
for _, s := range enabled {
assert.Equal(t, uint(10), s.AssistantID)
assert.True(t, s.Enabled)
}
}