217 lines
6.7 KiB
Go
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)
|
|
}
|
|
} |