Files
gochat/backend/internal/service/ai_takeover_test.go
T
2026-08-13 16:24:39 +08:00

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