H-150: stabilize PostgreSQL repository fixtures (#42)

Co-authored-by: Rogee <rogee@ipao.vip>
This commit is contained in:
Rogee
2026-08-18 00:57:44 +08:00
committed by GitHub
co-authored by rogee
parent a3b01e664e
commit dae6301a13
10 changed files with 134 additions and 63 deletions
@@ -249,10 +249,12 @@ func TestInboxAssignmentPolicyRepo_FindInboxesByPolicy(t *testing.T) {
policy := &model.AssignmentPolicy{AccountID: 1, Name: "P1"}
require.NoError(t, policyRepo.Create(ctx, policy))
require.NoError(t, db.Create(&model.Inbox{AccountID: 1, Name: "I1", ChannelType: "web_widget"}).Error)
require.NoError(t, db.Create(&model.Inbox{AccountID: 1, Name: "I2", ChannelType: "web_widget"}).Error)
require.NoError(t, repo.Create(ctx, &model.InboxAssignmentPolicy{InboxID: 1, AssignmentPolicyID: policy.ID}))
require.NoError(t, repo.Create(ctx, &model.InboxAssignmentPolicy{InboxID: 2, AssignmentPolicyID: policy.ID}))
inbox1 := &model.Inbox{AccountID: 1, Name: "I1", ChannelType: "web_widget"}
inbox2 := &model.Inbox{AccountID: 1, Name: "I2", ChannelType: "web_widget"}
require.NoError(t, db.Create(inbox1).Error)
require.NoError(t, db.Create(inbox2).Error)
require.NoError(t, repo.Create(ctx, &model.InboxAssignmentPolicy{InboxID: inbox1.ID, AssignmentPolicyID: policy.ID}))
require.NoError(t, repo.Create(ctx, &model.InboxAssignmentPolicy{InboxID: inbox2.ID, AssignmentPolicyID: policy.ID}))
inboxes, err := repo.FindInboxesByPolicy(ctx, 1, policy.ID)
require.NoError(t, err)
@@ -20,6 +20,7 @@ func createTestAttachment(t *testing.T, db *gorm.DB, messageID, accountID uint,
FileType: fileType,
FileName: fileName,
FileSize: 1024,
Metadata: `{}`,
}
require.NoError(t, db.Create(att).Error)
return att
@@ -109,6 +110,7 @@ func TestAttachmentRepo_Create_Success(t *testing.T) {
FileName: "video.mp4",
FileSize: 2048,
ExternalURL: "https://example.com/external",
Metadata: `{}`,
}
err := repo.Create(context.Background(), att)
@@ -196,4 +198,4 @@ func TestAttachmentRepo_DeleteByMessage_Success(t *testing.T) {
assert.NotNil(t, sd1.DeletedAt)
require.NoError(t, db.Unscoped().First(&sd2, att2.ID).Error)
assert.NotNil(t, sd2.DeletedAt)
}
}
@@ -7,10 +7,20 @@ import (
"github.com/pgvector/pgvector-go"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
"github.com/gochat/gochat/internal/model"
)
func createCaptainResponseTestAssistant(t *testing.T, db *gorm.DB) *model.CaptainAssistant {
t.Helper()
account := &model.Account{Name: "Captain response test", Active: true}
require.NoError(t, db.Create(account).Error)
assistant := &model.CaptainAssistant{AccountID: account.ID, Name: "Test assistant"}
require.NoError(t, db.Create(assistant).Error)
return assistant
}
// 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.
@@ -230,17 +240,18 @@ func TestCaptainAssistantResponseRepo_ListByDocument(t *testing.T) {
func TestCaptainAssistantResponseRepo_SimilaritySearch(t *testing.T) {
skipIfSQLite(t) // pgvector requires PostgreSQL
db := setupTestDB(t, &model.CaptainAssistantResponse{})
db := setupTestDB(t, &model.CaptainAssistant{}, &model.CaptainAssistantResponse{})
repo := NewCaptainAssistantResponseRepo(db)
assistant := createCaptainResponseTestAssistant(t, db)
// Create approved responses under assistant 10
// Create approved responses under the test assistant.
for i := 0; i < 3; i++ {
r := createTestResponse(1, 10, "SimilarQ-"+string(rune('A'+i)), "SimilarA-"+string(rune('A'+i)))
r := createTestResponse(assistant.AccountID, assistant.ID, "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 := createTestResponse(assistant.AccountID, assistant.ID, "RejectedQ", "RejectedA")
r.Status = model.ResponseStatusRejected
require.NoError(t, repo.Create(context.Background(), r))
@@ -251,13 +262,13 @@ func TestCaptainAssistantResponseRepo_SimilaritySearch(t *testing.T) {
}
embedding := pgvector.NewVector(dims)
results, err := repo.SimilaritySearch(context.Background(), 10, embedding, 5)
results, err := repo.SimilaritySearch(context.Background(), assistant.ID, 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)
assert.Equal(t, assistant.ID, resp.AssistantID)
}
}
@@ -266,12 +277,13 @@ func TestCaptainAssistantResponseRepo_SimilaritySearch(t *testing.T) {
func TestCaptainAssistantResponseRepo_SearchByEmbedding(t *testing.T) {
skipIfSQLite(t) // pgvector requires PostgreSQL
db := setupTestDB(t, &model.CaptainAssistantResponse{})
db := setupTestDB(t, &model.CaptainAssistant{}, &model.CaptainAssistantResponse{})
repo := NewCaptainAssistantResponseRepo(db)
assistant := createCaptainResponseTestAssistant(t, db)
// Create approved responses under assistant 20
// Create approved responses under the test assistant.
for i := 0; i < 2; i++ {
r := createTestResponse(1, 20, "EmbedQ-"+string(rune('A'+i)), "EmbedA-"+string(rune('A'+i)))
r := createTestResponse(assistant.AccountID, assistant.ID, "EmbedQ-"+string(rune('A'+i)), "EmbedA-"+string(rune('A'+i)))
r.Status = model.ResponseStatusApproved
require.NoError(t, repo.Create(context.Background(), r))
}
@@ -282,10 +294,10 @@ func TestCaptainAssistantResponseRepo_SearchByEmbedding(t *testing.T) {
}
embedding := pgvector.NewVector(dims)
results, err := repo.SearchByEmbedding(context.Background(), 20, embedding, 5)
results, err := repo.SearchByEmbedding(context.Background(), assistant.ID, embedding, 5)
require.NoError(t, err)
assert.Len(t, results, 2)
for _, resp := range results {
assert.Equal(t, uint(20), resp.AssistantID)
assert.Equal(t, assistant.ID, resp.AssistantID)
}
}
}
@@ -70,7 +70,7 @@ func TestCaptainCustomToolRepo_Create_WithBearerAuth(t *testing.T) {
// Verify auth_config persisted correctly
found, err := repo.GetByID(context.Background(), tool.ID)
require.NoError(t, err)
assert.Equal(t, tool.AuthConfig, found.AuthConfig)
assert.JSONEq(t, string(tool.AuthConfig), string(found.AuthConfig))
}
// ========== GetByID ==========
@@ -89,7 +89,7 @@ func TestCaptainCustomToolRepo_GetByID(t *testing.T) {
assert.Equal(t, tool.Title, found.Title)
assert.Equal(t, tool.Slug, found.Slug)
assert.Equal(t, tool.EndpointURL, found.EndpointURL)
assert.Equal(t, tool.ParamSchema, found.ParamSchema)
assert.JSONEq(t, string(tool.ParamSchema), string(found.ParamSchema))
}
func TestCaptainCustomToolRepo_GetByID_NotFound(t *testing.T) {
@@ -245,4 +245,4 @@ func TestCaptainCustomToolRepo_CountByAccount(t *testing.T) {
cnt, err = repo.CountByAccount(context.Background(), accountID)
require.NoError(t, err)
assert.Equal(t, int64(2), cnt)
}
}
@@ -40,7 +40,13 @@ func TestCaptainInboxRepo_Create_OnlyOneActiveAssistantPerInbox(t *testing.T) {
ctx := context.Background()
require.NoError(t, repo.Create(ctx, newTestCaptainInbox(1, 10, 100)))
if db.Dialector.Name() == "postgres" {
require.NoError(t, db.SavePoint("before_duplicate").Error)
}
require.Error(t, repo.Create(ctx, newTestCaptainInbox(2, 10, 100)))
if db.Dialector.Name() == "postgres" {
require.NoError(t, db.RollbackTo("before_duplicate").Error)
}
require.NoError(t, repo.Delete(ctx, 1, 10))
require.NoError(t, repo.Create(ctx, newTestCaptainInbox(2, 10, 100)))
}
@@ -198,8 +198,8 @@ func TestCaptainScenarioRepo_FindEnabled(t *testing.T) {
// Create 1 disabled scenario
s3 := createTestScenario(1, 10, "DisabledC")
s3.Enabled = false
require.NoError(t, repo.Create(context.Background(), s3))
require.NoError(t, db.Model(s3).Update("enabled", false).Error)
// Create 1 enabled scenario under different assistant
s4 := createTestScenario(1, 20, "OtherEnabled")
@@ -214,4 +214,4 @@ func TestCaptainScenarioRepo_FindEnabled(t *testing.T) {
assert.Equal(t, uint(10), s.AssistantID)
assert.True(t, s.Enabled)
}
}
}
@@ -408,7 +408,7 @@ func TestContactRepo_DeleteCustomAttributes(t *testing.T) {
// Verify custom attributes are set
found, err := repo.FindByID(context.Background(), contact.ID)
require.NoError(t, err)
assert.Equal(t, datatypes.JSON(`{"key":"value"}`), found.CustomAttributes)
assert.JSONEq(t, `{"key":"value"}`, string(found.CustomAttributes))
// Delete custom attributes
err = repo.DeleteCustomAttributes(context.Background(), contact.ID)
@@ -417,7 +417,7 @@ func TestContactRepo_DeleteCustomAttributes(t *testing.T) {
// Verify custom attributes are now empty
found2, err := repo.FindByID(context.Background(), contact.ID)
require.NoError(t, err)
assert.Equal(t, datatypes.JSON(`{}`), found2.CustomAttributes)
assert.JSONEq(t, `{}`, string(found2.CustomAttributes))
}
func TestContactRepo_DeleteCustomAttributes_NotFound(t *testing.T) {
@@ -19,30 +19,34 @@ func TestConversationRepo_Search_ILIKE(t *testing.T) {
db := setupTestDB(t)
repo := NewConversationRepo(db)
ctx := context.Background()
accountID := uint(1)
account := createTestAccountForSearch(t, db)
otherAccount := createTestAccountForSearch(t, db)
inbox := createTestInboxForSearch(t, db, account.ID, "PG search")
contact := createTestContactForSearch(t, db, account.ID, "PG contact", "pg-search@example.com", "")
otherContact := createTestContactForSearch(t, db, otherAccount.ID, "Other PG contact", "other-pg-search@example.com", "")
convs := []model.Conversation{
{AccountID: accountID, InboxID: 1, Status: "open", Labels: "support,billing"},
{AccountID: accountID, InboxID: 1, Status: "open", Labels: "billing,urgent"},
{AccountID: accountID, InboxID: 1, Status: "open", Labels: "sales,lead"},
{AccountID: accountID, InboxID: 1, Status: "resolved", Labels: "support,closed"},
{AccountID: uint(2), InboxID: 1, Status: "open", Labels: "support,other"},
{AccountID: account.ID, InboxID: inbox.ID, ContactID: contact.ID, Status: "open", Labels: "support,billing"},
{AccountID: account.ID, InboxID: inbox.ID, ContactID: contact.ID, Status: "open", Labels: "billing,urgent"},
{AccountID: account.ID, InboxID: inbox.ID, ContactID: contact.ID, Status: "open", Labels: "sales,lead"},
{AccountID: account.ID, InboxID: inbox.ID, ContactID: contact.ID, Status: "resolved", Labels: "support,closed"},
{AccountID: otherAccount.ID, InboxID: inbox.ID, ContactID: otherContact.ID, Status: "open", Labels: "support,other"},
}
for _, c := range convs {
require.NoError(t, repo.Create(ctx, &c))
}
results, total, err := repo.Search(ctx, accountID, "billing", 0, 10, search.SearchModeILike)
results, total, err := repo.Search(ctx, account.ID, "billing", 0, 10, search.SearchModeILike)
require.NoError(t, err)
assert.Equal(t, int64(2), total)
_ = results
results, total, err = repo.Search(ctx, accountID, "support", 0, 10, search.SearchModeILike)
results, total, err = repo.Search(ctx, account.ID, "support", 0, 10, search.SearchModeILike)
require.NoError(t, err)
assert.Equal(t, int64(2), total)
_ = results
results, total, err = repo.Search(ctx, accountID, "nonexistent", 0, 10, search.SearchModeILike)
results, total, err = repo.Search(ctx, account.ID, "nonexistent", 0, 10, search.SearchModeILike)
require.NoError(t, err)
assert.Equal(t, int64(0), total)
assert.Empty(t, results)
@@ -54,20 +58,20 @@ func TestConversationRepo_Search_Trigram(t *testing.T) {
db := setupTestDB(t)
repo := NewConversationRepo(db)
ctx := context.Background()
accountID := uint(1)
account := createTestAccountForSearch(t, db)
inbox := createTestInboxForSearch(t, db, account.ID, "PG trigram")
contact := createTestContactForSearch(t, db, account.ID, "PG trigram contact", "pg-trigram@example.com", "")
convs := []model.Conversation{
{AccountID: accountID, InboxID: 1, Status: "open", Labels: "customer support query"},
{AccountID: accountID, InboxID: 1, Status: "open", Labels: "billing inquiry"},
{AccountID: account.ID, InboxID: inbox.ID, ContactID: contact.ID, Status: "open", Labels: "customer"},
{AccountID: account.ID, InboxID: inbox.ID, ContactID: contact.ID, Status: "open", Labels: "billing inquiry"},
}
for _, c := range convs {
require.NoError(t, repo.Create(ctx, &c))
}
results, total, err := repo.Search(ctx, accountID, "custmr", 0, 10, search.SearchModeTrigram)
if err != nil {
t.Skipf("pg_trgm extension not available: %v", err)
}
results, total, err := repo.Search(ctx, account.ID, "custmr", 0, 10, search.SearchModeTrigram)
require.NoError(t, err)
assert.GreaterOrEqual(t, total, int64(1))
assert.NotEmpty(t, results)
}
@@ -172,4 +176,4 @@ func TestConversationRepo_UpdatedWithin_PG(t *testing.T) {
Where("account_id = ? AND updated_at > NOW() - INTERVAL '0 seconds'", accountID).
Count(&total)
assert.LessOrEqual(t, total, int64(1))
}
}
@@ -168,9 +168,9 @@ func TestReportingEventRepo_FindByDateRange(t *testing.T) {
assert.Equal(t, e1.ID, events[0].ID)
}
// --- Test 7: AggregateByMetric (PG-only: uses COALESCE/AVG aggregation) ---
// --- Test 7: AggregateByMetric (PG-only: uses COALESCE/SUM aggregation) ---
func TestReportingEventRepo_AggregateByMetric(t *testing.T) {
skipIfSQLite(t) // COALESCE(AVG(value), 0) aggregation uses PG-specific semantics
skipIfSQLite(t) // COALESCE(SUM(value), 0) aggregation uses PG-specific semantics
db := setupTestDB(t, &model.ReportingEvent{})
repo := NewReportingEventRepo(db)
@@ -191,10 +191,10 @@ func TestReportingEventRepo_AggregateByMetric(t *testing.T) {
e3.CreatedAt = time.Date(2025, 5, 10, 0, 0, 0, 0, time.UTC)
require.NoError(t, db.Create(e3).Error)
avgValue, total, err := repo.AggregateByMetric(context.Background(), 1, model.MetricNameFirstResponse, since, until)
sumValue, total, err := repo.AggregateByMetric(context.Background(), 1, model.MetricNameFirstResponse, since, until)
require.NoError(t, err)
assert.Equal(t, int64(3), total)
assert.InDelta(t, 20.0, avgValue, 0.01) // (10+20+30)/3 = 20
assert.InDelta(t, 60.0, sumValue, 0.01)
}
// --- Test 8: AggregateByMetric returns zero when no matches (PG-only) ---
@@ -207,8 +207,8 @@ func TestReportingEventRepo_AggregateByMetric_NoMatches(t *testing.T) {
since := time.Date(2025, 1, 1, 0, 0, 0, 0, time.UTC)
until := time.Date(2025, 12, 31, 23, 59, 59, 0, time.UTC)
avgValue, total, err := repo.AggregateByMetric(context.Background(), 999, model.MetricNameFirstResponse, since, until)
sumValue, total, err := repo.AggregateByMetric(context.Background(), 999, model.MetricNameFirstResponse, since, until)
require.NoError(t, err)
assert.Equal(t, int64(0), total)
assert.Equal(t, 0.0, avgValue)
}
assert.Equal(t, 0.0, sumValue)
}
+60 -15
View File
@@ -1,6 +1,7 @@
package repository
import (
"fmt"
"os"
"sync"
"testing"
@@ -21,6 +22,9 @@ var (
actualDBIsPG bool
actualDBChecked bool
actualDBCheckMu sync.Mutex
pgSchemaMu sync.Mutex
pgMigrated = make(map[string]bool)
pgExtensionsSet bool
)
// wantPostgres returns true when GOCHAT_TEST_DB != "sqlite".
@@ -78,24 +82,28 @@ func setupTestDB(t *testing.T, models ...interface{}) *gorm.DB {
allModels := append(defaultTestModels(), models...)
if isActuallyPG() {
db, err := gorm.Open(postgres.Open(pgDSN()), &gorm.Config{})
db, err := gorm.Open(postgres.Open(pgDSN()), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
Logger: logger.Default.LogMode(logger.Silent),
})
if err != nil {
t.Fatalf("failed to connect to PostgreSQL test db: %v\nSet GOCHAT_TEST_DB=sqlite to use SQLite instead", err)
}
if err := db.AutoMigrate(allModels...); err != nil {
t.Fatalf("failed to auto-migrate PG: %v", err)
if err := migratePGModels(db, allModels); err != nil {
t.Fatalf("failed to prepare PostgreSQL test schema: %v", err)
}
tx := db.Begin()
if tx.Error != nil {
t.Fatalf("failed to begin PostgreSQL test transaction: %v", tx.Error)
}
// Cleanup: truncate all tables after test
t.Cleanup(func() {
for _, m := range allModels {
db.Unscoped().Session(&gorm.Session{AllowGlobalUpdate: true}).Delete(m)
}
_ = tx.Rollback().Error
sqlDB, _ := db.DB()
sqlDB.Close()
})
return db
return tx
}
return openSQLiteDB(t, allModels)
@@ -109,29 +117,66 @@ func setupBenchmarkDB(b *testing.B, models ...interface{}) *gorm.DB {
if isActuallyPG() {
db, err := gorm.Open(postgres.Open(pgDSN()), &gorm.Config{
Logger: logger.Default.LogMode(logger.Silent),
DisableForeignKeyConstraintWhenMigrating: true,
Logger: logger.Default.LogMode(logger.Silent),
})
if err != nil {
b.Fatalf("failed to connect to PostgreSQL benchmark db: %v", err)
}
if err := db.AutoMigrate(allModels...); err != nil {
b.Fatalf("failed to auto-migrate PG: %v", err)
if err := migratePGModels(db, allModels); err != nil {
b.Fatalf("failed to prepare PostgreSQL benchmark schema: %v", err)
}
tx := db.Begin()
if tx.Error != nil {
b.Fatalf("failed to begin PostgreSQL benchmark transaction: %v", tx.Error)
}
b.Cleanup(func() {
for _, m := range allModels {
db.Unscoped().Session(&gorm.Session{AllowGlobalUpdate: true}).Delete(m)
}
_ = tx.Rollback().Error
sqlDB, _ := db.DB()
sqlDB.Close()
})
return db
return tx
}
return openSQLiteDBBench(b, allModels)
}
func migratePGModels(db *gorm.DB, models []interface{}) error {
pgSchemaMu.Lock()
defer pgSchemaMu.Unlock()
if !pgExtensionsSet {
if err := db.Exec("CREATE EXTENSION IF NOT EXISTS vector; CREATE EXTENSION IF NOT EXISTS pg_trgm").Error; err != nil {
return fmt.Errorf("enable extensions: %w", err)
}
pgExtensionsSet = true
}
pending := make([]interface{}, 0, len(models))
keys := make([]string, 0, len(models))
for _, m := range models {
key := fmt.Sprintf("%T", m)
if pgMigrated[key] {
continue
}
pgMigrated[key] = true
pending = append(pending, m)
keys = append(keys, key)
}
if len(pending) == 0 {
return nil
}
if err := db.AutoMigrate(pending...); err != nil {
for _, key := range keys {
delete(pgMigrated, key)
}
return fmt.Errorf("auto-migrate: %w", err)
}
return nil
}
// defaultTestModels returns the standard set of models needed for most tests.
func defaultTestModels() []interface{} {
return []interface{}{