272 lines
8.6 KiB
Go
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"])
|
|
} |