226 lines
7.9 KiB
Go
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)
|
|
}
|