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) } // --- 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) }