291 lines
9.9 KiB
Go
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)
|
|
}
|
|
} |