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

272 lines
8.6 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"
)
// createTestAssistant is a helper that builds a CaptainAssistant with typical fields.
func createTestAssistant(accountID uint, name string) *model.CaptainAssistant {
config, _ := json.Marshal(map[string]interface{}{
"temperature": 0.7,
"product_name": "TestBot",
})
guardrails, _ := json.Marshal(map[string]interface{}{
"max_tokens": 500,
"forbidden_topics": []string{"politics"},
})
responseGuidelines, _ := json.Marshal(map[string]interface{}{
"tone": "professional",
"length": "concise",
})
return &model.CaptainAssistant{
AccountID: accountID,
Name: name,
Description: "A test assistant",
Config: config,
Guardrails: guardrails,
ResponseGuidelines: responseGuidelines,
Status: model.AssistantStatusActive,
}
}
// --- 1. Create ---
func TestCaptainAssistantRepo_Create(t *testing.T) {
db := setupTestDB(t, &model.CaptainAssistant{})
repo := NewCaptainAssistantRepo(db)
assistant := createTestAssistant(1, "TestAssistant")
err := repo.Create(context.Background(), assistant)
require.NoError(t, err)
assert.NotZero(t, assistant.ID, "ID should be set after Create")
assert.Equal(t, "TestAssistant", assistant.Name)
assert.Equal(t, uint(1), assistant.AccountID)
}
// --- 2. GetByID ---
func TestCaptainAssistantRepo_GetByID(t *testing.T) {
db := setupTestDB(t, &model.CaptainAssistant{})
repo := NewCaptainAssistantRepo(db)
assistant := createTestAssistant(1, "GetByIDAssistant")
err := repo.Create(context.Background(), assistant)
require.NoError(t, err)
found, err := repo.GetByID(context.Background(), assistant.ID)
require.NoError(t, err)
assert.Equal(t, assistant.ID, found.ID)
assert.Equal(t, "GetByIDAssistant", found.Name)
assert.Equal(t, uint(1), found.AccountID)
}
// --- 3. GetByID Not Found ---
func TestCaptainAssistantRepo_GetByID_NotFound(t *testing.T) {
db := setupTestDB(t, &model.CaptainAssistant{})
repo := NewCaptainAssistantRepo(db)
found, err := repo.GetByID(context.Background(), 9999)
assert.Error(t, err)
assert.Nil(t, found)
}
// --- 4. Update ---
func TestCaptainAssistantRepo_Update(t *testing.T) {
db := setupTestDB(t, &model.CaptainAssistant{})
repo := NewCaptainAssistantRepo(db)
assistant := createTestAssistant(1, "BeforeUpdate")
err := repo.Create(context.Background(), assistant)
require.NoError(t, err)
assistant.Name = "AfterUpdate"
assistant.Description = "Updated description"
newConfig, _ := json.Marshal(map[string]interface{}{
"temperature": 0.9,
"product_name": "UpdatedBot",
})
assistant.Config = newConfig
err = repo.Update(context.Background(), assistant)
require.NoError(t, err)
found, err := repo.GetByID(context.Background(), assistant.ID)
require.NoError(t, err)
assert.Equal(t, "AfterUpdate", found.Name)
assert.Equal(t, "Updated description", found.Description)
// Verify Config JSON was updated
var foundConfig map[string]interface{}
require.NoError(t, json.Unmarshal(found.Config, &foundConfig))
assert.Equal(t, 0.9, foundConfig["temperature"])
assert.Equal(t, "UpdatedBot", foundConfig["product_name"])
}
// --- 5. Delete ---
func TestCaptainAssistantRepo_Delete(t *testing.T) {
db := setupTestDB(t, &model.CaptainAssistant{})
repo := NewCaptainAssistantRepo(db)
assistant := createTestAssistant(1, "DeleteMe")
err := repo.Create(context.Background(), assistant)
require.NoError(t, err)
err = repo.Delete(context.Background(), assistant.ID)
require.NoError(t, err)
// After deletion, GetByID should return error
found, err := repo.GetByID(context.Background(), assistant.ID)
assert.Error(t, err)
assert.Nil(t, found)
}
// --- 6. ListByAccount ---
func TestCaptainAssistantRepo_ListByAccount(t *testing.T) {
db := setupTestDB(t, &model.CaptainAssistant{})
repo := NewCaptainAssistantRepo(db)
// Create 3 assistants under account 1
for i := 0; i < 3; i++ {
a := createTestAssistant(1, "ListAssistant"+string(rune('A'+i)))
require.NoError(t, repo.Create(context.Background(), a))
}
// Create 1 assistant under account 2
a2 := createTestAssistant(2, "OtherAccountAssistant")
require.NoError(t, repo.Create(context.Background(), a2))
assistants, count, err := repo.ListByAccount(context.Background(), 1, 0, 10)
require.NoError(t, err)
assert.Equal(t, int64(3), count)
assert.Len(t, assistants, 3)
for _, a := range assistants {
assert.Equal(t, uint(1), a.AccountID)
}
}
// --- 7. ListByAccount with pagination ---
func TestCaptainAssistantRepo_ListByAccount_Pagination(t *testing.T) {
db := setupTestDB(t, &model.CaptainAssistant{})
repo := NewCaptainAssistantRepo(db)
// Create 5 assistants under account 1
for i := 0; i < 5; i++ {
a := createTestAssistant(1, "PagAssistant"+string(rune('A'+i)))
require.NoError(t, repo.Create(context.Background(), a))
}
// Fetch first page: offset=0, limit=2
assistants, count, err := repo.ListByAccount(context.Background(), 1, 0, 2)
require.NoError(t, err)
assert.Equal(t, int64(5), count, "total count should be 5 regardless of pagination")
assert.Len(t, assistants, 2, "first page should return 2 items")
// Fetch second page: offset=2, limit=2
assistants2, count2, err := repo.ListByAccount(context.Background(), 1, 2, 2)
require.NoError(t, err)
assert.Equal(t, int64(5), count2)
assert.Len(t, assistants2, 2, "second page should return 2 items")
}
// --- 8. FindByInboxID ---
func TestCaptainAssistantRepo_FindByInboxID(t *testing.T) {
skipIfSQLite(t) // JOIN on captain_inboxes uses PG-specific migration
db := setupTestDB(t, &model.CaptainAssistant{}, &model.CaptainInbox{})
repo := NewCaptainAssistantRepo(db)
// Create an assistant
assistant := createTestAssistant(1, "InboxLinkedAssistant")
err := repo.Create(context.Background(), assistant)
require.NoError(t, err)
// Create an inbox (need Account + Inbox first)
account := &model.Account{Name: "InboxTestAccount", Locale: "en"}
require.NoError(t, db.Create(account).Error)
inbox := &model.Inbox{Name: "TestInbox", AccountID: account.ID, ChannelType: "web"}
require.NoError(t, db.Create(inbox).Error)
// Link assistant to inbox via captain_inboxes
ci := &model.CaptainInbox{
AssistantID: assistant.ID,
InboxID: inbox.ID,
AccountID: account.ID,
}
require.NoError(t, db.Create(ci).Error)
// FindByInboxID should return the linked assistant
assistants, err := repo.FindByInboxID(context.Background(), inbox.ID)
require.NoError(t, err)
assert.Len(t, assistants, 1)
assert.Equal(t, assistant.ID, assistants[0].ID)
assert.Equal(t, "InboxLinkedAssistant", assistants[0].Name)
}
// --- 9. FindByInboxID no results ---
func TestCaptainAssistantRepo_FindByInboxID_NoResults(t *testing.T) {
skipIfSQLite(t)
db := setupTestDB(t, &model.CaptainAssistant{}, &model.CaptainInbox{})
repo := NewCaptainAssistantRepo(db)
// Query an inbox that has no linked assistants
assistants, err := repo.FindByInboxID(context.Background(), 9999)
require.NoError(t, err)
assert.Len(t, assistants, 0)
}
// --- 10. JSON fields (Config, Guardrails, ResponseGuidelines) round-trip ---
func TestCaptainAssistantRepo_JSONFields_RoundTrip(t *testing.T) {
db := setupTestDB(t, &model.CaptainAssistant{})
repo := NewCaptainAssistantRepo(db)
config, _ := json.Marshal(map[string]interface{}{
"temperature": 0.5,
"product_name": "RoundTripBot",
"feature_faq": true,
})
guardrails, _ := json.Marshal(map[string]interface{}{
"max_tokens": 1000,
"forbidden_topics": []string{"security", "finance"},
})
responseGuidelines, _ := json.Marshal(map[string]interface{}{
"tone": "friendly",
"language": "zh-CN",
})
assistant := &model.CaptainAssistant{
AccountID: 1,
Name: "JSONBot",
Config: config,
Guardrails: guardrails,
ResponseGuidelines: responseGuidelines,
Status: model.AssistantStatusActive,
}
require.NoError(t, repo.Create(context.Background(), assistant))
found, err := repo.GetByID(context.Background(), assistant.ID)
require.NoError(t, err)
// Verify Config
var cfg map[string]interface{}
require.NoError(t, json.Unmarshal(found.Config, &cfg))
assert.Equal(t, 0.5, cfg["temperature"])
assert.Equal(t, "RoundTripBot", cfg["product_name"])
// Verify Guardrails
var gr map[string]interface{}
require.NoError(t, json.Unmarshal(found.Guardrails, &gr))
assert.Equal(t, 1000.0, gr["max_tokens"])
// Verify ResponseGuidelines
var rg map[string]interface{}
require.NoError(t, json.Unmarshal(found.ResponseGuidelines, &rg))
assert.Equal(t, "friendly", rg["tone"])
assert.Equal(t, "zh-CN", rg["language"])
}