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

291 lines
9.9 KiB
Go

package repository
import (
"context"
"testing"
"github.com/pgvector/pgvector-go"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/gochat/gochat/internal/model"
)
// zeroEmbedding returns a 1536-dimensional zero vector suitable for SQLite compatibility.
// SQLite cannot scan an empty string as pgvector; providing an explicit zero vector
// ensures the embedding column is always populated with a parseable value.
func zeroEmbedding() pgvector.Vector {
dims := make([]float32, 1536)
return pgvector.NewVector(dims)
}
// createTestResponse builds a minimal valid CaptainAssistantResponse with a zero embedding.
func createTestResponse(accountID, assistantID uint, question, answer string) *model.CaptainAssistantResponse {
return &model.CaptainAssistantResponse{
AccountID: accountID,
AssistantID: assistantID,
Question: question,
Answer: answer,
Status: model.ResponseStatusApproved,
Edited: false,
Embedding: zeroEmbedding(),
}
}
// createTestResponseWithDocument builds a response linked to a documentable with a zero embedding.
func createTestResponseWithDocument(accountID, assistantID uint, documentableID uint, documentableType, question, answer string) *model.CaptainAssistantResponse {
resp := createTestResponse(accountID, assistantID, question, answer)
resp.DocumentableID = &documentableID
resp.DocumentableType = documentableType
return resp
}
// ========== 1. Create ==========
func TestCaptainAssistantResponseRepo_Create(t *testing.T) {
skipIfSQLite(t)
db := setupTestDB(t, &model.CaptainAssistantResponse{})
repo := NewCaptainAssistantResponseRepo(db)
resp := createTestResponse(1, 10, "What is GoChat?", "GoChat is an open-source chat platform.")
err := repo.Create(context.Background(), resp)
require.NoError(t, err)
assert.NotZero(t, resp.ID, "ID should be set after Create")
assert.Equal(t, uint(1), resp.AccountID)
assert.Equal(t, uint(10), resp.AssistantID)
assert.Equal(t, "What is GoChat?", resp.Question)
assert.Equal(t, "GoChat is an open-source chat platform.", resp.Answer)
assert.Equal(t, model.ResponseStatusApproved, resp.Status)
}
// ========== 2. GetByID ==========
func TestCaptainAssistantResponseRepo_GetByID(t *testing.T) {
skipIfSQLite(t)
db := setupTestDB(t, &model.CaptainAssistantResponse{})
repo := NewCaptainAssistantResponseRepo(db)
resp := createTestResponse(1, 10, "GetByID question", "GetByID answer")
err := repo.Create(context.Background(), resp)
require.NoError(t, err)
found, err := repo.GetByID(context.Background(), resp.ID)
require.NoError(t, err)
assert.Equal(t, resp.ID, found.ID)
assert.Equal(t, "GetByID question", found.Question)
assert.Equal(t, uint(10), found.AssistantID)
}
// ========== 3. GetByID Not Found ==========
func TestCaptainAssistantResponseRepo_GetByID_NotFound(t *testing.T) {
skipIfSQLite(t)
db := setupTestDB(t, &model.CaptainAssistantResponse{})
repo := NewCaptainAssistantResponseRepo(db)
found, err := repo.GetByID(context.Background(), 9999)
assert.Error(t, err)
assert.Nil(t, found)
}
// ========== 4. Update ==========
func TestCaptainAssistantResponseRepo_Update(t *testing.T) {
skipIfSQLite(t)
db := setupTestDB(t, &model.CaptainAssistantResponse{})
repo := NewCaptainAssistantResponseRepo(db)
resp := createTestResponse(1, 10, "Original question", "Original answer")
err := repo.Create(context.Background(), resp)
require.NoError(t, err)
// Modify fields and update
resp.Question = "Updated question"
resp.Answer = "Updated answer"
resp.Status = model.ResponseStatusPending
resp.Edited = true
err = repo.Update(context.Background(), resp)
require.NoError(t, err)
found, err := repo.GetByID(context.Background(), resp.ID)
require.NoError(t, err)
assert.Equal(t, "Updated question", found.Question)
assert.Equal(t, "Updated answer", found.Answer)
assert.Equal(t, model.ResponseStatusPending, found.Status)
assert.True(t, found.Edited)
}
// ========== 5. Delete ==========
func TestCaptainAssistantResponseRepo_Delete(t *testing.T) {
skipIfSQLite(t)
db := setupTestDB(t, &model.CaptainAssistantResponse{})
repo := NewCaptainAssistantResponseRepo(db)
resp := createTestResponse(1, 10, "Delete question", "Delete answer")
err := repo.Create(context.Background(), resp)
require.NoError(t, err)
err = repo.Delete(context.Background(), resp.ID)
require.NoError(t, err)
// After deletion, GetByID should return error
found, err := repo.GetByID(context.Background(), resp.ID)
assert.Error(t, err)
assert.Nil(t, found)
}
// ========== 6. ListByAssistant ==========
func TestCaptainAssistantResponseRepo_ListByAssistant(t *testing.T) {
skipIfSQLite(t)
db := setupTestDB(t, &model.CaptainAssistantResponse{})
repo := NewCaptainAssistantResponseRepo(db)
// Create 3 responses under assistant 10, account 1
for i := 0; i < 3; i++ {
r := createTestResponse(1, 10, "Q10-"+string(rune('A'+i)), "A10-"+string(rune('A'+i)))
require.NoError(t, repo.Create(context.Background(), r))
}
// Create 2 responses under assistant 20, account 1 (should not appear)
for i := 0; i < 2; i++ {
r := createTestResponse(1, 20, "Q20-"+string(rune('A'+i)), "A20-"+string(rune('A'+i)))
require.NoError(t, repo.Create(context.Background(), r))
}
responses, count, err := repo.ListByAssistant(context.Background(), 10, 0, 10)
require.NoError(t, err)
assert.Equal(t, int64(3), count)
assert.Len(t, responses, 3)
for _, r := range responses {
assert.Equal(t, uint(10), r.AssistantID)
}
}
// ========== 7. ListByAssistant with Pagination ==========
func TestCaptainAssistantResponseRepo_ListByAssistant_Pagination(t *testing.T) {
skipIfSQLite(t)
db := setupTestDB(t, &model.CaptainAssistantResponse{})
repo := NewCaptainAssistantResponseRepo(db)
// Create 5 responses under assistant 10
for i := 0; i < 5; i++ {
r := createTestResponse(1, 10, "PagQ-"+string(rune('A'+i)), "PagA-"+string(rune('A'+i)))
require.NoError(t, repo.Create(context.Background(), r))
}
// First page: offset=0, limit=2
responses, count, err := repo.ListByAssistant(context.Background(), 10, 0, 2)
require.NoError(t, err)
assert.Equal(t, int64(5), count, "total count should be 5 regardless of pagination")
assert.Len(t, responses, 2, "first page should return 2 items")
// Second page: offset=2, limit=2
responses2, count2, err := repo.ListByAssistant(context.Background(), 10, 2, 2)
require.NoError(t, err)
assert.Equal(t, int64(5), count2)
assert.Len(t, responses2, 2, "second page should return 2 items")
// Third page: offset=4, limit=2
responses3, count3, err := repo.ListByAssistant(context.Background(), 10, 4, 2)
require.NoError(t, err)
assert.Equal(t, int64(5), count3)
assert.Len(t, responses3, 1, "last page should return 1 item")
}
// ========== 8. ListByDocument ==========
func TestCaptainAssistantResponseRepo_ListByDocument(t *testing.T) {
skipIfSQLite(t)
db := setupTestDB(t, &model.CaptainAssistantResponse{})
repo := NewCaptainAssistantResponseRepo(db)
// Create 2 responses linked to document (ID=100, type="CaptainDocument")
for i := 0; i < 2; i++ {
r := createTestResponseWithDocument(1, 10, 100, "CaptainDocument",
"DocQ-"+string(rune('A'+i)), "DocA-"+string(rune('A'+i)))
require.NoError(t, repo.Create(context.Background(), r))
}
// Create 1 response linked to assistant (documentableType="CaptainAssistant") — should not appear
r := createTestResponseWithDocument(1, 10, 50, "CaptainAssistant",
"AssistQ", "AssistA")
require.NoError(t, repo.Create(context.Background(), r))
responses, count, err := repo.ListByDocument(context.Background(), 100, "CaptainDocument", 0, 10)
require.NoError(t, err)
assert.Equal(t, int64(2), count)
assert.Len(t, responses, 2)
for _, resp := range responses {
assert.Equal(t, uint(100), *resp.DocumentableID)
assert.Equal(t, "CaptainDocument", resp.DocumentableType)
}
}
// ========== 9. SimilaritySearch (PG only) ==========
func TestCaptainAssistantResponseRepo_SimilaritySearch(t *testing.T) {
skipIfSQLite(t) // pgvector requires PostgreSQL
db := setupTestDB(t, &model.CaptainAssistantResponse{})
repo := NewCaptainAssistantResponseRepo(db)
// Create approved responses under assistant 10
for i := 0; i < 3; i++ {
r := createTestResponse(1, 10, "SimilarQ-"+string(rune('A'+i)), "SimilarA-"+string(rune('A'+i)))
r.Status = model.ResponseStatusApproved
require.NoError(t, repo.Create(context.Background(), r))
}
// Create one rejected response — should be excluded from similarity search
r := createTestResponse(1, 10, "RejectedQ", "RejectedA")
r.Status = model.ResponseStatusRejected
require.NoError(t, repo.Create(context.Background(), r))
// Build a dummy 1536-dimensional embedding vector for search
dims := make([]float32, 1536)
for i := range dims {
dims[i] = 0.01
}
embedding := pgvector.NewVector(dims)
results, err := repo.SimilaritySearch(context.Background(), 10, embedding, 5)
require.NoError(t, err)
// Only approved responses (3) should be returned
assert.Len(t, results, 3)
for _, resp := range results {
assert.Equal(t, model.ResponseStatusApproved, resp.Status)
assert.Equal(t, uint(10), resp.AssistantID)
}
}
// ========== 10. SearchByEmbedding (alias for SimilaritySearch, PG only) ==========
func TestCaptainAssistantResponseRepo_SearchByEmbedding(t *testing.T) {
skipIfSQLite(t) // pgvector requires PostgreSQL
db := setupTestDB(t, &model.CaptainAssistantResponse{})
repo := NewCaptainAssistantResponseRepo(db)
// Create approved responses under assistant 20
for i := 0; i < 2; i++ {
r := createTestResponse(1, 20, "EmbedQ-"+string(rune('A'+i)), "EmbedA-"+string(rune('A'+i)))
r.Status = model.ResponseStatusApproved
require.NoError(t, repo.Create(context.Background(), r))
}
dims := make([]float32, 1536)
for i := range dims {
dims[i] = 0.02
}
embedding := pgvector.NewVector(dims)
results, err := repo.SearchByEmbedding(context.Background(), 20, embedding, 5)
require.NoError(t, err)
assert.Len(t, results, 2)
for _, resp := range results {
assert.Equal(t, uint(20), resp.AssistantID)
}
}