199 lines
11 KiB
Go
199 lines
11 KiB
Go
package service
|
|
|
|
import (
|
|
"context"
|
|
"errors"
|
|
"fmt"
|
|
"testing"
|
|
|
|
"github.com/gochat/gochat/internal/channel"
|
|
"github.com/gochat/gochat/internal/model"
|
|
"github.com/gochat/gochat/internal/worker"
|
|
"github.com/stretchr/testify/assert"
|
|
"github.com/stretchr/testify/require"
|
|
"gorm.io/gorm"
|
|
)
|
|
|
|
func configureInboxAI(t *testing.T, svc *ConversationService, accountID, inboxID uint) (*model.AgentBot, *model.CaptainAssistant) {
|
|
t.Helper()
|
|
db := svc.DB()
|
|
require.NoError(t, db.AutoMigrate(&model.AgentBot{}, &model.AgentBotInbox{}, &model.CaptainAssistant{}, &model.CaptainInbox{}))
|
|
assistant := &model.CaptainAssistant{AccountID: accountID, Name: "Channel AI", Status: model.AssistantStatusActive, Config: []byte(`{}`)}
|
|
require.NoError(t, db.Create(assistant).Error)
|
|
require.NoError(t, db.Create(&model.CaptainInbox{AccountID: accountID, InboxID: inboxID, AssistantID: assistant.ID}).Error)
|
|
bot := &model.AgentBot{AccountID: &accountID, Name: "Channel AI", BotType: "captain", Secret: fmt.Sprintf("secret-%d", assistant.ID), AccessToken: fmt.Sprintf("token-%d", assistant.ID), Config: []byte(fmt.Sprintf(`{"assistant_id":%d}`, assistant.ID))}
|
|
require.NoError(t, db.Create(bot).Error)
|
|
require.NoError(t, db.Create(&model.AgentBotInbox{AgentBotID: bot.ID, InboxID: inboxID, Status: model.AgentBotInboxActive}).Error)
|
|
return bot, assistant
|
|
}
|
|
|
|
func TestConversationServiceAITakeoverStartExitAndRestart(t *testing.T) {
|
|
svc, db := setupConversationService(t)
|
|
capture := &captureConversationEventsListener{}
|
|
svc.dispatcher.Register(capture)
|
|
account := createConversationServiceTestAccount(t, db)
|
|
inbox := createConversationServiceTestInbox(t, db, account.ID)
|
|
contact := createConversationServiceTestContact(t, db, account.ID)
|
|
conversation := createConversationServiceTestConversation(t, db, account.ID, inbox.ID, contact.ID, string(model.ConversationStatusOpen))
|
|
bot, _ := configureInboxAI(t, svc, account.ID, inbox.ID)
|
|
|
|
started, err := svc.StartAITakeover(context.Background(), account.ID, conversation.ID)
|
|
require.NoError(t, err)
|
|
require.NotNil(t, started.AssigneeAgentBotID)
|
|
assert.Equal(t, bot.ID, *started.AssigneeAgentBotID)
|
|
assert.Equal(t, string(model.ConversationStatusPending), started.Status)
|
|
assert.Equal(t, uint(1), started.AITakeoverVersion)
|
|
|
|
repeated, err := svc.StartAITakeover(context.Background(), account.ID, conversation.ID)
|
|
require.NoError(t, err)
|
|
assert.Equal(t, uint(1), repeated.AITakeoverVersion)
|
|
|
|
exited, err := svc.ExitAITakeover(context.Background(), account.ID, conversation.ID)
|
|
require.NoError(t, err)
|
|
assert.Nil(t, exited.AssigneeAgentBotID)
|
|
assert.Equal(t, string(model.ConversationStatusOpen), exited.Status)
|
|
assert.Equal(t, uint(2), exited.AITakeoverVersion)
|
|
require.GreaterOrEqual(t, len(capture.events), 2)
|
|
assert.Equal(t, channel.EventConversationUpdated, capture.events[len(capture.events)-2].Type)
|
|
assert.Equal(t, channel.EventConversationUnassigned, capture.events[len(capture.events)-1].Type)
|
|
|
|
eventCount := len(capture.events)
|
|
repeatedExit, err := svc.ExitAITakeover(context.Background(), account.ID, conversation.ID)
|
|
require.NoError(t, err)
|
|
assert.Nil(t, repeatedExit.AssigneeAgentBotID)
|
|
assert.Equal(t, eventCount, len(capture.events))
|
|
|
|
restarted, err := svc.StartAITakeover(context.Background(), account.ID, conversation.ID)
|
|
require.NoError(t, err)
|
|
assert.Equal(t, uint(3), restarted.AITakeoverVersion)
|
|
}
|
|
|
|
func TestConversationServiceAITakeoverRejectsMissingOrAmbiguousAI(t *testing.T) {
|
|
svc, db := setupConversationService(t)
|
|
account := createConversationServiceTestAccount(t, db)
|
|
inbox := createConversationServiceTestInbox(t, db, account.ID)
|
|
contact := createConversationServiceTestContact(t, db, account.ID)
|
|
conversation := createConversationServiceTestConversation(t, db, account.ID, inbox.ID, contact.ID, string(model.ConversationStatusOpen))
|
|
require.NoError(t, db.AutoMigrate(&model.AgentBot{}, &model.AgentBotInbox{}, &model.CaptainAssistant{}, &model.CaptainInbox{}))
|
|
|
|
_, err := svc.StartAITakeover(context.Background(), account.ID, conversation.ID)
|
|
assert.ErrorContains(t, err, "exactly one active AI")
|
|
|
|
bot, _ := configureInboxAI(t, svc, account.ID, inbox.ID)
|
|
otherBot := &model.AgentBot{AccountID: &account.ID, Name: "Other Captain", BotType: "captain", Secret: "other-secret", AccessToken: "other-token", Config: []byte(`{"assistant_id":999}`)}
|
|
require.NoError(t, db.Create(otherBot).Error)
|
|
require.NoError(t, db.Create(&model.AgentBotInbox{AgentBotID: otherBot.ID, InboxID: inbox.ID, Status: model.AgentBotInboxActive}).Error)
|
|
require.NotEqual(t, bot.ID, otherBot.ID)
|
|
_, err = svc.StartAITakeover(context.Background(), account.ID, conversation.ID)
|
|
assert.ErrorContains(t, err, "exactly one active AI")
|
|
}
|
|
|
|
func TestConversationServiceAITakeoverRejectsConflictWithoutCreatingBinding(t *testing.T) {
|
|
svc, db := setupConversationService(t)
|
|
account := createConversationServiceTestAccount(t, db)
|
|
inbox := createConversationServiceTestInbox(t, db, account.ID)
|
|
contact := createConversationServiceTestContact(t, db, account.ID)
|
|
conversation := createConversationServiceTestConversation(t, db, account.ID, inbox.ID, contact.ID, string(model.ConversationStatusOpen))
|
|
require.NoError(t, db.AutoMigrate(&model.AgentBot{}, &model.AgentBotInbox{}, &model.CaptainAssistant{}, &model.CaptainInbox{}))
|
|
assistant := &model.CaptainAssistant{AccountID: account.ID, Name: "Captain", Status: model.AssistantStatusActive}
|
|
require.NoError(t, db.Create(assistant).Error)
|
|
require.NoError(t, db.Create(&model.CaptainInbox{AccountID: account.ID, InboxID: inbox.ID, AssistantID: assistant.ID}).Error)
|
|
conflictingBot := &model.AgentBot{AccountID: &account.ID, Name: "Webhook", BotType: "webhook"}
|
|
require.NoError(t, db.Create(conflictingBot).Error)
|
|
require.NoError(t, db.Create(&model.AgentBotInbox{AgentBotID: conflictingBot.ID, InboxID: inbox.ID, Status: model.AgentBotInboxActive}).Error)
|
|
|
|
_, err := svc.StartAITakeover(context.Background(), account.ID, conversation.ID)
|
|
require.ErrorContains(t, err, "exactly one active AI")
|
|
var captainBotCount int64
|
|
require.NoError(t, db.Model(&model.AgentBot{}).Where("bot_type = ?", "captain").Count(&captainBotCount).Error)
|
|
assert.Zero(t, captainBotCount)
|
|
var bindingCount int64
|
|
require.NoError(t, db.Model(&model.AgentBotInbox{}).Count(&bindingCount).Error)
|
|
assert.Equal(t, int64(1), bindingCount)
|
|
require.NoError(t, db.First(conversation, conversation.ID).Error)
|
|
assert.Equal(t, string(model.ConversationStatusOpen), conversation.Status)
|
|
assert.Nil(t, conversation.AssigneeAgentBotID)
|
|
assert.Zero(t, conversation.AITakeoverVersion)
|
|
}
|
|
|
|
func TestConversationServiceAITakeoverRollsBackBindingWhenTakeoverFails(t *testing.T) {
|
|
svc, db := setupConversationService(t)
|
|
account := createConversationServiceTestAccount(t, db)
|
|
inbox := createConversationServiceTestInbox(t, db, account.ID)
|
|
contact := createConversationServiceTestContact(t, db, account.ID)
|
|
conversation := createConversationServiceTestConversation(t, db, account.ID, inbox.ID, contact.ID, string(model.ConversationStatusOpen))
|
|
require.NoError(t, db.AutoMigrate(&model.AgentBot{}, &model.AgentBotInbox{}, &model.CaptainAssistant{}, &model.CaptainInbox{}))
|
|
assistant := &model.CaptainAssistant{AccountID: account.ID, Name: "Captain", Status: model.AssistantStatusActive}
|
|
require.NoError(t, db.Create(assistant).Error)
|
|
require.NoError(t, db.Create(&model.CaptainInbox{AccountID: account.ID, InboxID: inbox.ID, AssistantID: assistant.ID}).Error)
|
|
require.NoError(t, db.Callback().Update().Before("gorm:update").Register("test:fail_takeover", func(tx *gorm.DB) {
|
|
if tx.Statement.Table == "conversations" {
|
|
tx.AddError(errors.New("takeover failed"))
|
|
}
|
|
}))
|
|
|
|
_, err := svc.StartAITakeover(context.Background(), account.ID, conversation.ID)
|
|
require.ErrorContains(t, err, "takeover failed")
|
|
var botCount, bindingCount int64
|
|
require.NoError(t, db.Model(&model.AgentBot{}).Count(&botCount).Error)
|
|
require.NoError(t, db.Model(&model.AgentBotInbox{}).Count(&bindingCount).Error)
|
|
assert.Zero(t, botCount)
|
|
assert.Zero(t, bindingCount)
|
|
}
|
|
|
|
func TestMessageServiceHumanOutgoingExitsAITakeoverAtomically(t *testing.T) {
|
|
db, _, dispatcher, svc := setupMessageServiceWithDefaultLLM(t)
|
|
capture := &captureConversationEventsListener{}
|
|
dispatcher.Register(capture)
|
|
account := createTestAccount(t, db)
|
|
user := createTestUser(t, db, account.ID)
|
|
inbox := createTestInbox(t, db, account.ID, "web_widget")
|
|
contact := createTestContact(t, db, account.ID)
|
|
conversation := createTestConversation(t, db, account.ID, inbox.ID, contact.ID)
|
|
botID := uint(91)
|
|
require.NoError(t, db.Model(conversation).Updates(map[string]any{
|
|
"assignee_agent_bot_id": botID, "status": model.ConversationStatusPending, "ai_takeover_version": 1,
|
|
}).Error)
|
|
|
|
message, err := svc.Create(context.Background(), account.ID, user.ID, CreateMessageRequest{
|
|
ConversationID: conversation.ID, Content: "human reply", MessageType: string(model.MessageTypeOutgoing), ContentType: string(model.MessageContentTypeText),
|
|
})
|
|
require.NoError(t, err)
|
|
assert.True(t, message.AITakeoverExited)
|
|
|
|
var updated model.Conversation
|
|
require.NoError(t, db.First(&updated, conversation.ID).Error)
|
|
assert.Nil(t, updated.AssigneeAgentBotID)
|
|
assert.Equal(t, string(model.ConversationStatusOpen), updated.Status)
|
|
assert.Equal(t, uint(2), updated.AITakeoverVersion)
|
|
require.GreaterOrEqual(t, len(capture.events), 2)
|
|
assert.Equal(t, channel.EventConversationUpdated, capture.events[0].Type)
|
|
eventConversation := capture.events[0].Data["conversation"].(*model.Conversation)
|
|
assert.Nil(t, eventConversation.AssigneeAgentBotID)
|
|
assert.Equal(t, string(model.ConversationStatusOpen), eventConversation.Status)
|
|
assert.Equal(t, channel.EventMessageCreated, capture.events[1].Type)
|
|
}
|
|
|
|
func TestCaptainConversationResponseQueuesOncePerIncomingMessageAndSupportsMultipleTurns(t *testing.T) {
|
|
db, conversationSvc, messageSvc, account, _, conversation, _ := setupCaptainConversationWorkerTest(t)
|
|
wp := worker.NewWorkerPool(db)
|
|
conversationSvc.SetWorkerPool(wp)
|
|
messageSvc.SetWorkerPool(wp)
|
|
|
|
first, err := messageSvc.Create(context.Background(), account.ID, 99, CreateMessageRequest{ConversationID: conversation.ID, Content: "one", MessageType: string(model.MessageTypeIncoming), ContentType: string(model.MessageContentTypeText)})
|
|
require.NoError(t, err)
|
|
second, err := messageSvc.Create(context.Background(), account.ID, 99, CreateMessageRequest{ConversationID: conversation.ID, Content: "two", MessageType: string(model.MessageTypeIncoming), ContentType: string(model.MessageContentTypeText)})
|
|
require.NoError(t, err)
|
|
|
|
_, err = EnqueueCaptainConversationResponseForMessage(context.Background(), wp, db, first.ID)
|
|
require.NoError(t, err)
|
|
var count int64
|
|
require.NoError(t, db.Model(&model.BackgroundJob{}).Where("job_type = ?", TaskTypeCaptainConversationResponseBuilder).Count(&count).Error)
|
|
assert.Equal(t, int64(2), count)
|
|
|
|
var jobs []model.BackgroundJob
|
|
require.NoError(t, db.Where("job_type = ?", TaskTypeCaptainConversationResponseBuilder).Order("id").Find(&jobs).Error)
|
|
assert.Contains(t, jobs[0].IdempotencyKey, fmt.Sprintf("message:%d", first.ID))
|
|
assert.Contains(t, jobs[1].IdempotencyKey, fmt.Sprintf("message:%d", second.ID))
|
|
}
|