Files
gochat/internal/repository/copilot_message_repo_test.go
T

226 lines
7.9 KiB
Go

package repository
import (
"context"
"encoding/json"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
"github.com/gochat/gochat/internal/model"
)
// createTestCopilotMessage is a helper that builds a CopilotMessage with typical fields.
func createTestCopilotMessage(accountID, threadID uint, messageType model.CopilotMessageType, content string) *model.CopilotMessage {
message, _ := json.Marshal(map[string]interface{}{
"content": content,
})
return &model.CopilotMessage{
AccountID: accountID,
CopilotThreadID: threadID,
MessageType: messageType,
Message: message,
}
}
// createTestCopilotThread is a helper that creates a CopilotThread in the DB.
func createTestCopilotThread(t *testing.T, db *gorm.DB, accountID, userID uint, title string) *model.CopilotThread {
t.Helper()
thread := &model.CopilotThread{
AccountID: accountID,
UserID: userID,
Title: title,
}
require.NoError(t, db.Create(thread).Error)
return thread
}
// --- 1. Create ---
func TestCopilotMessageRepo_Create(t *testing.T) {
db := setupTestDB(t, &model.CopilotThread{}, &model.CopilotMessage{})
repo := NewCopilotMessageRepo(db)
thread := createTestCopilotThread(t, db, 1, 1, "Test Thread")
msg := createTestCopilotMessage(1, thread.ID, model.CopilotMessageTypeUser, "Hello copilot")
err := repo.Create(context.Background(), msg)
require.NoError(t, err)
assert.NotZero(t, msg.ID, "ID should be set after Create")
assert.Equal(t, uint(1), msg.AccountID)
assert.Equal(t, thread.ID, msg.CopilotThreadID)
assert.Equal(t, model.CopilotMessageTypeUser, msg.MessageType)
}
func TestCopilotMessageRepo_Create_AllowsChatwootMessageKeys(t *testing.T) {
db := setupTestDB(t, &model.CopilotThread{}, &model.CopilotMessage{})
repo := NewCopilotMessageRepo(db)
thread := createTestCopilotThread(t, db, 1, 1, "Tool Thread")
message, _ := json.Marshal(map[string]interface{}{
"content": "Looking up docs",
"reasoning": "Need product data",
"function_name": "search_documentation",
"reply_suggestion": "Try this reply",
})
err := repo.Create(context.Background(), &model.CopilotMessage{
AccountID: 1,
CopilotThreadID: thread.ID,
MessageType: model.CopilotMessageTypeAssistant,
Message: message,
})
require.NoError(t, err)
}
func TestCopilotMessageRepo_Create_RejectsUnknownMessageKeys(t *testing.T) {
db := setupTestDB(t, &model.CopilotThread{}, &model.CopilotMessage{})
repo := NewCopilotMessageRepo(db)
thread := createTestCopilotThread(t, db, 1, 1, "Bad Tool Thread")
message, _ := json.Marshal(map[string]interface{}{
"content": "hello",
"unknown": "bad",
})
err := repo.Create(context.Background(), &model.CopilotMessage{
AccountID: 1,
CopilotThreadID: thread.ID,
MessageType: model.CopilotMessageTypeAssistant,
Message: message,
})
require.Error(t, err)
assert.Contains(t, err.Error(), "invalid attribute: unknown")
}
// --- 2. GetByID ---
func TestCopilotMessageRepo_GetByID(t *testing.T) {
db := setupTestDB(t, &model.CopilotThread{}, &model.CopilotMessage{})
repo := NewCopilotMessageRepo(db)
thread := createTestCopilotThread(t, db, 1, 1, "Test Thread")
msg := createTestCopilotMessage(1, thread.ID, model.CopilotMessageTypeAssistant, "Hi there")
require.NoError(t, db.Create(msg).Error)
found, err := repo.GetByID(context.Background(), msg.ID)
require.NoError(t, err)
assert.Equal(t, msg.ID, found.ID)
assert.Equal(t, msg.CopilotThreadID, found.CopilotThreadID)
assert.Equal(t, model.CopilotMessageTypeAssistant, found.MessageType)
}
func TestCopilotMessageRepo_GetByID_NotFound(t *testing.T) {
db := setupTestDB(t, &model.CopilotThread{}, &model.CopilotMessage{})
repo := NewCopilotMessageRepo(db)
found, err := repo.GetByID(context.Background(), 9999)
assert.Error(t, err)
assert.Nil(t, found)
}
// --- 3. FindByThreadID ---
func TestCopilotMessageRepo_FindByThreadID(t *testing.T) {
db := setupTestDB(t, &model.CopilotThread{}, &model.CopilotMessage{})
repo := NewCopilotMessageRepo(db)
thread := createTestCopilotThread(t, db, 1, 1, "Test Thread")
// Create multiple messages in the thread
msg1 := createTestCopilotMessage(1, thread.ID, model.CopilotMessageTypeUser, "First message")
msg2 := createTestCopilotMessage(1, thread.ID, model.CopilotMessageTypeAssistant, "Second message")
msg3 := createTestCopilotMessage(1, thread.ID, model.CopilotMessageTypeAssistantThinking, "Thinking...")
require.NoError(t, db.Create(msg1).Error)
require.NoError(t, db.Create(msg2).Error)
require.NoError(t, db.Create(msg3).Error)
msgs, count, err := repo.FindByThreadID(context.Background(), thread.ID, 0, 10)
require.NoError(t, err)
assert.Equal(t, int64(3), count)
assert.Len(t, msgs, 3)
// Messages should be ordered by created_at ASC
assert.Equal(t, msg1.ID, msgs[0].ID)
assert.Equal(t, msg2.ID, msgs[1].ID)
assert.Equal(t, msg3.ID, msgs[2].ID)
}
func TestCopilotMessageRepo_FindByThreadID_Pagination(t *testing.T) {
db := setupTestDB(t, &model.CopilotThread{}, &model.CopilotMessage{})
repo := NewCopilotMessageRepo(db)
thread := createTestCopilotThread(t, db, 1, 1, "Paginated Thread")
// Create 5 messages
for i := 0; i < 5; i++ {
msg := createTestCopilotMessage(1, thread.ID, model.CopilotMessageTypeUser, "Message "+string(rune('A'+i)))
require.NoError(t, db.Create(msg).Error)
}
// Fetch first page (offset=0, limit=2)
msgs, count, err := repo.FindByThreadID(context.Background(), thread.ID, 0, 2)
require.NoError(t, err)
assert.Equal(t, int64(5), count)
assert.Len(t, msgs, 2)
// Fetch second page (offset=2, limit=2)
msgs2, count2, err := repo.FindByThreadID(context.Background(), thread.ID, 2, 2)
require.NoError(t, err)
assert.Equal(t, int64(5), count2)
assert.Len(t, msgs2, 2)
assert.NotEqual(t, msgs[0].ID, msgs2[0].ID, "different pages should return different messages")
}
// --- 4. ListByThread (alias for FindByThreadID) ---
func TestCopilotMessageRepo_ListByThread(t *testing.T) {
db := setupTestDB(t, &model.CopilotThread{}, &model.CopilotMessage{})
repo := NewCopilotMessageRepo(db)
thread := createTestCopilotThread(t, db, 1, 1, "ListByThread Thread")
msg1 := createTestCopilotMessage(1, thread.ID, model.CopilotMessageTypeUser, "User msg")
msg2 := createTestCopilotMessage(1, thread.ID, model.CopilotMessageTypeAssistant, "Assistant reply")
require.NoError(t, db.Create(msg1).Error)
require.NoError(t, db.Create(msg2).Error)
msgs, count, err := repo.ListByThread(context.Background(), thread.ID, 0, 10)
require.NoError(t, err)
assert.Equal(t, int64(2), count)
assert.Len(t, msgs, 2)
// Verify ListByThread returns the same results as FindByThreadID
msgsFind, countFind, errFind := repo.FindByThreadID(context.Background(), thread.ID, 0, 10)
require.NoError(t, errFind)
assert.Equal(t, count, countFind)
assert.Equal(t, msgs[0].ID, msgsFind[0].ID)
assert.Equal(t, msgs[1].ID, msgsFind[1].ID)
}
// --- 5. DeleteByThread ---
func TestCopilotMessageRepo_DeleteByThread(t *testing.T) {
db := setupTestDB(t, &model.CopilotThread{}, &model.CopilotMessage{})
repo := NewCopilotMessageRepo(db)
thread := createTestCopilotThread(t, db, 1, 1, "Delete Thread")
msg1 := createTestCopilotMessage(1, thread.ID, model.CopilotMessageTypeUser, "To be deleted")
msg2 := createTestCopilotMessage(1, thread.ID, model.CopilotMessageTypeAssistant, "Also deleted")
require.NoError(t, db.Create(msg1).Error)
require.NoError(t, db.Create(msg2).Error)
// Verify messages exist before deletion
msgs, count, err := repo.FindByThreadID(context.Background(), thread.ID, 0, 10)
require.NoError(t, err)
assert.Equal(t, int64(2), count)
assert.Len(t, msgs, 2)
// Delete all messages in the thread
err = repo.DeleteByThread(context.Background(), thread.ID)
require.NoError(t, err)
// Verify messages are gone
msgsAfter, countAfter, errAfter := repo.FindByThreadID(context.Background(), thread.ID, 0, 10)
require.NoError(t, errAfter)
assert.Equal(t, int64(0), countAfter)
assert.Len(t, msgsAfter, 0)
}