H-337: restore Captain inbox takeover and KB citations (#65)
* fix(captain): restore inbox takeover and KB citations * fix(captain): harden grounded citations and smoke seed --------- Co-authored-by: Rogee <rogee@ipao.vip>
This commit is contained in:
@@ -0,0 +1,114 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"github.com/gochat/gochat/internal/model"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/driver/sqlite"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/logger"
|
||||
)
|
||||
|
||||
func captainSeedTestDB(t *testing.T) *gorm.DB {
|
||||
t.Helper()
|
||||
dsn := filepath.Join(t.TempDir(), "captain-seed.db") + "?_busy_timeout=5000&_journal_mode=WAL"
|
||||
db, err := gorm.Open(sqlite.Open(dsn), &gorm.Config{
|
||||
DisableForeignKeyConstraintWhenMigrating: true,
|
||||
Logger: logger.Default.LogMode(logger.Silent),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, db.AutoMigrate(&model.Account{}, &model.CaptainAssistant{}, &model.CaptainInbox{}))
|
||||
return db
|
||||
}
|
||||
|
||||
func TestSeedSmokeCaptainIsRepeatableAndConcurrent(t *testing.T) {
|
||||
db := captainSeedTestDB(t)
|
||||
ctx := context.Background()
|
||||
account := &model.Account{Name: "Smoke account", Active: true}
|
||||
require.NoError(t, db.Create(account).Error)
|
||||
|
||||
start := make(chan struct{})
|
||||
errs := make(chan error, 2)
|
||||
var wg sync.WaitGroup
|
||||
for range 2 {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
<-start
|
||||
_, err := seedSmokeCaptain(ctx, db, account.ID, 10)
|
||||
errs <- err
|
||||
}()
|
||||
}
|
||||
close(start)
|
||||
wg.Wait()
|
||||
close(errs)
|
||||
for err := range errs {
|
||||
require.NoError(t, err)
|
||||
}
|
||||
first, err := seedSmokeCaptain(ctx, db, account.ID, 10)
|
||||
require.NoError(t, err)
|
||||
second, err := seedSmokeCaptain(ctx, db, account.ID, 10)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, first.ID, second.ID)
|
||||
|
||||
var assistants, bindings int64
|
||||
require.NoError(t, db.Model(&model.CaptainAssistant{}).Where("account_id = ? AND name = ?", account.ID, "Smoke Captain").Count(&assistants).Error)
|
||||
require.NoError(t, db.Model(&model.CaptainInbox{}).Where("inbox_id = ?", 10).Count(&bindings).Error)
|
||||
assert.Equal(t, int64(1), assistants)
|
||||
assert.Equal(t, int64(1), bindings)
|
||||
|
||||
var persisted model.CaptainAssistant
|
||||
require.NoError(t, db.First(&persisted, first.ID).Error)
|
||||
var config map[string]any
|
||||
require.NoError(t, json.Unmarshal(persisted.Config, &config))
|
||||
assert.Equal(t, true, config["feature_citation"])
|
||||
}
|
||||
|
||||
func TestSeedSmokeCaptainConflictHasNoSideEffects(t *testing.T) {
|
||||
db := captainSeedTestDB(t)
|
||||
account := &model.Account{Name: "Smoke account", Active: true}
|
||||
require.NoError(t, db.Create(account).Error)
|
||||
other := &model.CaptainAssistant{AccountID: account.ID, Name: "Other Captain", Status: model.AssistantStatusActive, Config: json.RawMessage(`{"feature_citation":true}`)}
|
||||
require.NoError(t, db.Create(other).Error)
|
||||
require.NoError(t, db.Create(&model.CaptainInbox{AccountID: account.ID, AssistantID: other.ID, InboxID: 10}).Error)
|
||||
smoke := &model.CaptainAssistant{AccountID: account.ID, Name: "Smoke Captain", Status: model.AssistantStatusActive, Config: json.RawMessage(`{"feature_citation":false,"sentinel":"keep"}`)}
|
||||
require.NoError(t, db.Create(smoke).Error)
|
||||
|
||||
_, err := seedSmokeCaptain(context.Background(), db, account.ID, 10)
|
||||
require.Error(t, err)
|
||||
|
||||
var persisted model.CaptainAssistant
|
||||
require.NoError(t, db.First(&persisted, smoke.ID).Error)
|
||||
var config map[string]any
|
||||
require.NoError(t, json.Unmarshal(persisted.Config, &config))
|
||||
assert.Equal(t, false, config["feature_citation"])
|
||||
assert.Equal(t, "keep", config["sentinel"])
|
||||
|
||||
var bindings []model.CaptainInbox
|
||||
require.NoError(t, db.Where("inbox_id = ?", 10).Find(&bindings).Error)
|
||||
require.Len(t, bindings, 1)
|
||||
assert.Equal(t, other.ID, bindings[0].AssistantID)
|
||||
}
|
||||
|
||||
func TestSeedSmokeCaptainRechecksSameAssistantAfterConflict(t *testing.T) {
|
||||
db := captainSeedTestDB(t)
|
||||
account := &model.Account{Name: "Smoke account", Active: true}
|
||||
require.NoError(t, db.Create(account).Error)
|
||||
smoke := &model.CaptainAssistant{AccountID: account.ID, Name: "Smoke Captain", Status: model.AssistantStatusActive, Config: json.RawMessage(`{"feature_citation":false}`)}
|
||||
require.NoError(t, db.Create(smoke).Error)
|
||||
require.NoError(t, db.Create(&model.CaptainInbox{AccountID: account.ID, AssistantID: smoke.ID, InboxID: 10}).Error)
|
||||
|
||||
persisted, err := seedSmokeCaptain(context.Background(), db, account.ID, 10)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, smoke.ID, persisted.ID)
|
||||
|
||||
var config map[string]any
|
||||
require.NoError(t, json.Unmarshal(persisted.Config, &config))
|
||||
assert.Equal(t, true, config["feature_citation"])
|
||||
}
|
||||
@@ -8,6 +8,7 @@ import (
|
||||
"fmt"
|
||||
"os"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/gochat/gochat/internal/app"
|
||||
@@ -19,6 +20,7 @@ import (
|
||||
applogger "github.com/gochat/gochat/pkg/logger"
|
||||
"gorm.io/datatypes"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/clause"
|
||||
)
|
||||
|
||||
func main() {
|
||||
@@ -139,7 +141,19 @@ type smokeSeedSummary struct {
|
||||
AgentBotID uint `json:"agent_bot_id"`
|
||||
}
|
||||
|
||||
var smokeCaptainSeedMu sync.Mutex
|
||||
|
||||
func seedSmokeData(ctx context.Context, db *gorm.DB) (*smokeSeedSummary, error) {
|
||||
var summary *smokeSeedSummary
|
||||
err := db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
var err error
|
||||
summary, err = seedSmokeDataInTransaction(ctx, tx)
|
||||
return err
|
||||
})
|
||||
return summary, err
|
||||
}
|
||||
|
||||
func seedSmokeDataInTransaction(ctx context.Context, db *gorm.DB) (*smokeSeedSummary, error) {
|
||||
adminEmail := getenvDefault("GOCHAT_SEED_ADMIN_EMAIL", "admin@gochat.local")
|
||||
adminPassword := getenvDefault("GOCHAT_SEED_ADMIN_PASSWORD", "changeme")
|
||||
adminName := getenvDefault("GOCHAT_SEED_ADMIN_NAME", "Super Admin")
|
||||
@@ -281,9 +295,9 @@ func seedSmokeData(ctx context.Context, db *gorm.DB) (*smokeSeedSummary, error)
|
||||
return nil, fmt.Errorf("seed capacity limit: %w", err)
|
||||
}
|
||||
|
||||
assistant := &model.CaptainAssistant{}
|
||||
if err := db.WithContext(ctx).Where("account_id = ? AND name = ?", account.ID, "Smoke Captain").FirstOrCreate(assistant, model.CaptainAssistant{AccountID: account.ID, Name: "Smoke Captain", Description: "B12 smoke assistant", Status: model.AssistantStatusActive, Config: json.RawMessage(`{"model":"gpt-4o"}`)}).Error; err != nil {
|
||||
return nil, fmt.Errorf("seed captain assistant: %w", err)
|
||||
assistant, err := seedSmokeCaptain(ctx, db, account.ID, inbox.ID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
captainMessage := &model.Message{}
|
||||
if err := db.WithContext(ctx).Where("conversation_id = ? AND content = ?", conversation.ID, "Smoke Captain answer").FirstOrCreate(captainMessage, model.Message{AccountID: account.ID, InboxID: inbox.ID, ConversationID: conversation.ID, SenderID: &assistant.ID, SenderType: "Captain::Assistant", MessageType: "outgoing", ContentType: "text", Content: "Smoke Captain answer", Status: "sent"}).Error; err != nil {
|
||||
@@ -301,6 +315,61 @@ func seedSmokeData(ctx context.Context, db *gorm.DB) (*smokeSeedSummary, error)
|
||||
return &smokeSeedSummary{AdminID: admin.ID, AdminEmail: adminEmail, AdminPassword: adminPassword, AccountID: account.ID, InboxID: inbox.ID, VoiceInboxID: voiceInbox.ID, ContactID: contact.ID, CompanyID: company.ID, PortalID: portal.ID, ArticleID: article.ID, ConversationID: conversation.ID, ConversationDisplayID: conversationDisplayID, ConversationUID: conversation.UUID, CsatMessageID: csatMessage.ID, SlaPolicyID: sla.ID, CustomRoleID: customRole.ID, CapacityPolicyID: capacity.ID, CaptainAssistantID: assistant.ID, CaptainMessageID: captainMessage.ID, AgentBotID: agentBot.ID}, nil
|
||||
}
|
||||
|
||||
func seedSmokeCaptain(ctx context.Context, db *gorm.DB, accountID, inboxID uint) (*model.CaptainAssistant, error) {
|
||||
// ponytail: seed is an operator-only path; this lock keeps SQLite and
|
||||
// concurrent in-process runs deterministic, while PostgreSQL also locks the account row.
|
||||
smokeCaptainSeedMu.Lock()
|
||||
defer smokeCaptainSeedMu.Unlock()
|
||||
|
||||
assistant := &model.CaptainAssistant{}
|
||||
err := db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
var account model.Account
|
||||
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Select("id").First(&account, accountID).Error; err != nil {
|
||||
return fmt.Errorf("seed captain account: %w", err)
|
||||
}
|
||||
if err := tx.Where("account_id = ? AND name = ?", accountID, "Smoke Captain").Attrs(model.CaptainAssistant{AccountID: accountID, Name: "Smoke Captain", Description: "B12 smoke assistant", Status: model.AssistantStatusActive, Config: json.RawMessage(`{"model":"gpt-4o","feature_citation":true}`)}).FirstOrCreate(assistant).Error; err != nil {
|
||||
return fmt.Errorf("seed captain assistant: %w", err)
|
||||
}
|
||||
|
||||
binding := &model.CaptainInbox{AccountID: accountID, AssistantID: assistant.ID, InboxID: inboxID}
|
||||
if err := tx.Clauses(clause.OnConflict{
|
||||
Columns: []clause.Column{{Name: "inbox_id"}},
|
||||
TargetWhere: clause.Where{Exprs: []clause.Expression{clause.Expr{SQL: "deleted_at IS NULL"}}},
|
||||
DoNothing: true,
|
||||
}).Create(binding).Error; err != nil {
|
||||
return fmt.Errorf("seed captain inbox: %w", err)
|
||||
}
|
||||
var persisted model.CaptainInbox
|
||||
if err := tx.Where("inbox_id = ?", inboxID).First(&persisted).Error; err != nil {
|
||||
return fmt.Errorf("seed captain inbox: %w", err)
|
||||
}
|
||||
if persisted.AccountID != accountID || persisted.AssistantID != assistant.ID {
|
||||
return fmt.Errorf("seed captain inbox: inbox belongs to assistant %d in account %d, expected assistant %d in account %d", persisted.AssistantID, persisted.AccountID, assistant.ID, accountID)
|
||||
}
|
||||
|
||||
var config map[string]any
|
||||
if err := json.Unmarshal(assistant.Config, &config); err != nil {
|
||||
return fmt.Errorf("seed captain assistant config: %w", err)
|
||||
}
|
||||
if config == nil {
|
||||
config = make(map[string]any)
|
||||
}
|
||||
if citations, _ := config["feature_citation"].(bool); !citations {
|
||||
config["feature_citation"] = true
|
||||
configJSON, err := json.Marshal(config)
|
||||
if err != nil {
|
||||
return fmt.Errorf("seed captain assistant config: %w", err)
|
||||
}
|
||||
if err := tx.Model(assistant).Update("config", configJSON).Error; err != nil {
|
||||
return fmt.Errorf("seed captain assistant config: %w", err)
|
||||
}
|
||||
assistant.Config = configJSON
|
||||
}
|
||||
return nil
|
||||
})
|
||||
return assistant, err
|
||||
}
|
||||
|
||||
func seedMessage(ctx context.Context, db *gorm.DB, conversation *model.Conversation, inboxID, senderID uint, messageType, contentType, content string) (*model.Message, error) {
|
||||
message := &model.Message{}
|
||||
var count int64
|
||||
|
||||
@@ -214,6 +214,7 @@ func setupCaptainAssistantHandlerTestWithProvider(t *testing.T, provider llm.Pro
|
||||
&model.Inbox{},
|
||||
&model.CaptainAssistant{},
|
||||
&model.CaptainInbox{},
|
||||
&model.CaptainDocument{},
|
||||
&model.CaptainAssistantResponse{},
|
||||
))
|
||||
t.Cleanup(func() {
|
||||
@@ -534,7 +535,7 @@ func (p *captainPlaygroundFakeProvider) CreateEmbedding(ctx context.Context, req
|
||||
if p.embeddingErr != nil || p.embeddingResponseSet || p.embeddingResponse != nil {
|
||||
return p.embeddingResponse, p.embeddingErr
|
||||
}
|
||||
return &llm.EmbeddingResponse{}, nil
|
||||
return &llm.EmbeddingResponse{Data: []llm.EmbeddingData{{Embedding: []float64{0.1, 0.2, 0.3}}}}, nil
|
||||
}
|
||||
|
||||
func (p *captainPlaygroundFakeProvider) ChatCompletionStream(ctx context.Context, req llm.ChatRequest, onChunk func(llm.StreamChunk) error) error {
|
||||
|
||||
@@ -107,6 +107,7 @@ type AssistantConfig struct {
|
||||
TemperatureConfigured bool `json:"-"`
|
||||
FeatureFAQ bool `json:"feature_faq,omitempty"`
|
||||
FeatureMemory bool `json:"feature_memory,omitempty"`
|
||||
FeatureCitation bool `json:"feature_citation,omitempty"`
|
||||
FeatureContactAttributes bool `json:"feature_contact_attributes,omitempty"`
|
||||
ProductName string `json:"product_name,omitempty"`
|
||||
Instructions string `json:"instructions,omitempty"`
|
||||
@@ -225,6 +226,7 @@ type CaptainAssistantResponse struct {
|
||||
Edited bool `gorm:"default:false;not null" json:"edited"`
|
||||
// pgvector-go Vector type for embeddings (dimension follows the configured embedding model)
|
||||
Embedding pgvector.Vector `gorm:"type:vector" json:"embedding,omitempty"`
|
||||
SourceURL string `gorm:"column:source_url;->;-:migration" json:"-"`
|
||||
|
||||
Assistant CaptainAssistant `gorm:"foreignKey:AssistantID" json:"assistant,omitempty"`
|
||||
}
|
||||
|
||||
@@ -106,9 +106,11 @@ func (r *CaptainAssistantResponseRepo) SimilaritySearch(ctx context.Context, ass
|
||||
// Cosine distance (<=>) orders by closest first.
|
||||
// Omit embedding column from SELECT — the pgvector stub can't scan it back.
|
||||
if err := r.db.WithContext(ctx).
|
||||
Select("id, account_id, assistant_id, documentable_id, documentable_type, question, answer, status, edited, created_at, updated_at").
|
||||
Where("assistant_id = ? AND status = ?", assistantID, model.ResponseStatusApproved).
|
||||
Order(gorm.Expr("embedding <=> ?::vector", vecStr)).
|
||||
Select("captain_assistant_responses.id, captain_assistant_responses.account_id, captain_assistant_responses.assistant_id, captain_assistant_responses.documentable_id, captain_assistant_responses.documentable_type, captain_assistant_responses.question, captain_assistant_responses.answer, captain_assistant_responses.status, captain_assistant_responses.edited, captain_assistant_responses.created_at, captain_assistant_responses.updated_at, COALESCE(captain_documents.external_link, '') AS source_url").
|
||||
Joins("JOIN captain_assistants ON captain_assistants.id = captain_assistant_responses.assistant_id AND captain_assistants.account_id = captain_assistant_responses.account_id AND captain_assistants.deleted_at IS NULL").
|
||||
Joins("LEFT JOIN captain_documents ON captain_documents.id = captain_assistant_responses.documentable_id AND captain_assistant_responses.documentable_type IN ? AND captain_documents.account_id = captain_assistant_responses.account_id AND captain_documents.assistant_id = captain_assistant_responses.assistant_id AND captain_documents.deleted_at IS NULL", []string{"Captain::Document", "CaptainDocument"}).
|
||||
Where("captain_assistant_responses.assistant_id = ? AND captain_assistant_responses.status = ? AND captain_assistant_responses.embedding IS NOT NULL", assistantID, model.ResponseStatusApproved).
|
||||
Order(gorm.Expr("captain_assistant_responses.embedding <=> ?::vector", vecStr)).
|
||||
Limit(limit).
|
||||
Find(&responses).Error; err != nil {
|
||||
return nil, err
|
||||
|
||||
@@ -277,16 +277,41 @@ func TestCaptainAssistantResponseRepo_SimilaritySearch(t *testing.T) {
|
||||
func TestCaptainAssistantResponseRepo_SearchByEmbedding(t *testing.T) {
|
||||
skipIfSQLite(t) // pgvector requires PostgreSQL
|
||||
|
||||
db := setupTestDB(t, &model.CaptainAssistant{}, &model.CaptainAssistantResponse{})
|
||||
db := setupTestDB(t, &model.CaptainAssistant{}, &model.CaptainDocument{}, &model.CaptainAssistantResponse{})
|
||||
repo := NewCaptainAssistantResponseRepo(db)
|
||||
assistant := createCaptainResponseTestAssistant(t, db)
|
||||
otherAssistant := &model.CaptainAssistant{AccountID: assistant.AccountID, Name: "Other assistant"}
|
||||
require.NoError(t, db.Create(otherAssistant).Error)
|
||||
otherAccount := &model.Account{Name: "Other account", Active: true}
|
||||
require.NoError(t, db.Create(otherAccount).Error)
|
||||
foreignAssistant := &model.CaptainAssistant{AccountID: otherAccount.ID, Name: "Foreign assistant"}
|
||||
require.NoError(t, db.Create(foreignAssistant).Error)
|
||||
|
||||
// Create approved responses under the test assistant.
|
||||
for i := 0; i < 2; i++ {
|
||||
r := createTestResponse(assistant.AccountID, assistant.ID, "EmbedQ-"+string(rune('A'+i)), "EmbedA-"+string(rune('A'+i)))
|
||||
r.Status = model.ResponseStatusApproved
|
||||
require.NoError(t, repo.Create(context.Background(), r))
|
||||
document := &model.CaptainDocument{AccountID: assistant.AccountID, AssistantID: assistant.ID, Name: "Source", ExternalLink: "https://example.com/knowledge"}
|
||||
crossAssistantDocument := &model.CaptainDocument{AccountID: assistant.AccountID, AssistantID: otherAssistant.ID, Name: "Wrong assistant", ExternalLink: "https://example.com/wrong-assistant"}
|
||||
crossAccountDocument := &model.CaptainDocument{AccountID: otherAccount.ID, AssistantID: foreignAssistant.ID, Name: "Wrong account", ExternalLink: "https://example.com/wrong-account"}
|
||||
require.NoError(t, db.Create(document).Error)
|
||||
require.NoError(t, db.Create(crossAssistantDocument).Error)
|
||||
require.NoError(t, db.Create(crossAccountDocument).Error)
|
||||
|
||||
responses := []*model.CaptainAssistantResponse{
|
||||
createTestResponseWithDocument(assistant.AccountID, assistant.ID, document.ID, "Captain::Document", "Valid source", "Valid answer"),
|
||||
createTestResponse(assistant.AccountID, assistant.ID, "No source", "Compatible answer"),
|
||||
createTestResponseWithDocument(assistant.AccountID, assistant.ID, crossAssistantDocument.ID, "Captain::Document", "Wrong assistant source", "Must not cite"),
|
||||
createTestResponseWithDocument(assistant.AccountID, assistant.ID, crossAccountDocument.ID, "Captain::Document", "Wrong account source", "Must not cite"),
|
||||
}
|
||||
for _, response := range responses {
|
||||
require.NoError(t, repo.Create(context.Background(), response))
|
||||
}
|
||||
|
||||
nullEmbedding := createTestResponse(assistant.AccountID, assistant.ID, "Missing embedding", "Must not retrieve")
|
||||
require.NoError(t, repo.Create(context.Background(), nullEmbedding))
|
||||
require.NoError(t, db.Model(nullEmbedding).UpdateColumn("embedding", gorm.Expr("NULL")).Error)
|
||||
corruptAccount := createTestResponse(otherAccount.ID, assistant.ID, "Wrong response account", "Must not retrieve")
|
||||
require.NoError(t, repo.Create(context.Background(), corruptAccount))
|
||||
allNull := createTestResponse(otherAssistant.AccountID, otherAssistant.ID, "Only null", "Must not retrieve")
|
||||
require.NoError(t, repo.Create(context.Background(), allNull))
|
||||
require.NoError(t, db.Model(allNull).UpdateColumn("embedding", gorm.Expr("NULL")).Error)
|
||||
|
||||
dims := make([]float32, 1536)
|
||||
for i := range dims {
|
||||
@@ -294,10 +319,24 @@ func TestCaptainAssistantResponseRepo_SearchByEmbedding(t *testing.T) {
|
||||
}
|
||||
embedding := pgvector.NewVector(dims)
|
||||
|
||||
results, err := repo.SearchByEmbedding(context.Background(), assistant.ID, embedding, 5)
|
||||
results, err := repo.SearchByEmbedding(context.Background(), assistant.ID, embedding, 10)
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, results, 2)
|
||||
assert.Len(t, results, 4)
|
||||
for _, resp := range results {
|
||||
assert.Equal(t, assistant.ID, resp.AssistantID)
|
||||
assert.Equal(t, assistant.AccountID, resp.AccountID)
|
||||
assert.NotEqual(t, nullEmbedding.ID, resp.ID)
|
||||
assert.NotEqual(t, corruptAccount.ID, resp.ID)
|
||||
}
|
||||
sources := make([]string, 0, len(results))
|
||||
for _, result := range results {
|
||||
sources = append(sources, result.SourceURL)
|
||||
}
|
||||
assert.Contains(t, sources, document.ExternalLink)
|
||||
assert.NotContains(t, sources, crossAssistantDocument.ExternalLink)
|
||||
assert.NotContains(t, sources, crossAccountDocument.ExternalLink)
|
||||
|
||||
empty, err := repo.SearchByEmbedding(context.Background(), otherAssistant.ID, embedding, 5)
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, empty, "an assistant with only NULL embeddings must return no grounded results")
|
||||
}
|
||||
|
||||
@@ -35,12 +35,17 @@ func TestConversationServiceAITakeoverStartExitAndRestart(t *testing.T) {
|
||||
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)
|
||||
require.NoError(t, db.AutoMigrate(&model.AgentBot{}, &model.AgentBotInbox{}, &model.CaptainAssistant{}, &model.CaptainInbox{}))
|
||||
assistant := &model.CaptainAssistant{AccountID: account.ID, Name: "Channel AI", Status: model.AssistantStatusActive, Config: []byte(`{}`)}
|
||||
require.NoError(t, db.Create(assistant).Error)
|
||||
require.NoError(t, db.Create(&model.CaptainInbox{AccountID: account.ID, InboxID: inbox.ID, AssistantID: assistant.ID}).Error)
|
||||
|
||||
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)
|
||||
var bot model.AgentBot
|
||||
require.NoError(t, db.First(&bot, *started.AssigneeAgentBotID).Error)
|
||||
assert.Equal(t, assistant.ID, extractAssistantIDFromBotConfig(bot.Config))
|
||||
assert.Equal(t, string(model.ConversationStatusPending), started.Status)
|
||||
assert.Equal(t, uint(1), started.AITakeoverVersion)
|
||||
|
||||
|
||||
@@ -17,21 +17,25 @@ func TestCaptainAssistantFAQRetrievalSeparatesEmptyAndFailures(t *testing.T) {
|
||||
history := []PlaygroundMessage{{Role: "user", Content: "How do refunds work?"}}
|
||||
embedding := &llm.EmbeddingResponse{Data: []llm.EmbeddingData{{Embedding: []float64{0.1, 0.2, 0.3}}}}
|
||||
|
||||
t.Run("disabled FAQ skips retrieval and still allows chat", func(t *testing.T) {
|
||||
t.Run("conversation FAQ generation flag does not disable approved knowledge", func(t *testing.T) {
|
||||
provider := &mockLLMProvider{
|
||||
embeddingError: errors.New("must not be called"),
|
||||
chatResponse: &llm.ChatResponse{Choices: []llm.ChatChoice{{Message: llm.ChatMessage{Content: "Chat without FAQ."}}}},
|
||||
embeddingResponse: embedding,
|
||||
chatResponse: &llm.ChatResponse{Choices: []llm.ChatChoice{{Message: llm.ChatMessage{Content: "37 days [[1](https://example.com/knowledge)]."}}}},
|
||||
}
|
||||
repo := &mockResponseRepo{searchByEmbeddingError: errors.New("must not be called")}
|
||||
repo := &mockResponseRepo{searchByEmbeddingResult: []model.CaptainAssistantResponse{{
|
||||
Question: "What is the warranty window?", Answer: "37 days.", SourceURL: "https://example.com/knowledge",
|
||||
}}}
|
||||
svc := &CaptainAssistantService{responseRepo: repo, llmProvider: provider, promptBuilder: NewSystemPromptBuilder()}
|
||||
disabled := &model.CaptainAssistant{Name: "Fin", Status: model.AssistantStatusActive, Config: []byte(`{"feature_faq":false}`)}
|
||||
disabled := &model.CaptainAssistant{Name: "Fin", Status: model.AssistantStatusActive, Config: []byte(`{"feature_faq":false,"feature_citation":true}`)}
|
||||
|
||||
content, err := svc.generatePlaygroundLLMResponse(context.Background(), disabled, history)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "Chat without FAQ.", content)
|
||||
assert.Zero(t, provider.embeddingCalls)
|
||||
assert.Zero(t, repo.searchByEmbeddingCalls)
|
||||
assert.Equal(t, "37 days [[1](https://example.com/knowledge)].", content)
|
||||
assert.Equal(t, 1, provider.embeddingCalls)
|
||||
assert.Equal(t, 1, repo.searchByEmbeddingCalls)
|
||||
require.NotNil(t, provider.lastChatRequest)
|
||||
assert.Contains(t, provider.lastChatRequest.Messages[0].Content, "Source: https://example.com/knowledge")
|
||||
assert.Contains(t, provider.lastChatRequest.Messages[0].Content, "[[n](URL)]")
|
||||
})
|
||||
|
||||
t.Run("empty results still allow chat", func(t *testing.T) {
|
||||
@@ -39,15 +43,77 @@ func TestCaptainAssistantFAQRetrievalSeparatesEmptyAndFailures(t *testing.T) {
|
||||
embeddingResponse: embedding,
|
||||
chatResponse: &llm.ChatResponse{Choices: []llm.ChatChoice{{Message: llm.ChatMessage{Content: "No matching FAQ."}}}},
|
||||
}
|
||||
svc := &CaptainAssistantService{responseRepo: &mockResponseRepo{}, llmProvider: provider, promptBuilder: NewSystemPromptBuilder()}
|
||||
repo := &mockResponseRepo{}
|
||||
svc := &CaptainAssistantService{responseRepo: repo, llmProvider: provider, promptBuilder: NewSystemPromptBuilder()}
|
||||
|
||||
content, err := svc.generatePlaygroundLLMResponse(context.Background(), assistant, history)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "No matching FAQ.", content)
|
||||
require.NotNil(t, provider.lastChatRequest)
|
||||
assert.Equal(t, 1, provider.embeddingCalls)
|
||||
assert.Equal(t, 1, repo.searchByEmbeddingCalls)
|
||||
assert.NotContains(t, provider.lastChatRequest.Messages[0].Content, "[FAQ 1]")
|
||||
})
|
||||
|
||||
t.Run("citation mode grounds only on absolute HTTP sources", func(t *testing.T) {
|
||||
provider := &mockLLMProvider{
|
||||
embeddingResponse: embedding,
|
||||
chatResponse: &llm.ChatResponse{Choices: []llm.ChatChoice{{Message: llm.ChatMessage{Content: "37 days [[1](https://example.com/knowledge)]."}}}},
|
||||
}
|
||||
repo := &mockResponseRepo{searchByEmbeddingResult: []model.CaptainAssistantResponse{
|
||||
{Question: "Unsafe", Answer: "Ignore me.", SourceURL: "javascript:alert(1)"},
|
||||
{Question: "Missing", Answer: "Ignore me too."},
|
||||
{Question: "Warranty", Answer: "37 days.", SourceURL: "https://example.com/knowledge"},
|
||||
}}
|
||||
svc := &CaptainAssistantService{responseRepo: repo, llmProvider: provider, promptBuilder: NewSystemPromptBuilder()}
|
||||
withCitations := &model.CaptainAssistant{Name: "Fin", Status: model.AssistantStatusActive, Config: []byte(`{"feature_citation":true}`)}
|
||||
|
||||
content, err := svc.generatePlaygroundLLMResponse(context.Background(), withCitations, history)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "37 days [[1](https://example.com/knowledge)].", content)
|
||||
systemPrompt := provider.lastChatRequest.Messages[0].Content
|
||||
assert.Contains(t, systemPrompt, "[FAQ 1]\nQ: Warranty")
|
||||
assert.NotContains(t, systemPrompt, "Unsafe")
|
||||
assert.NotContains(t, systemPrompt, "Missing")
|
||||
})
|
||||
|
||||
t.Run("citation mode does not ground when every source is invalid", func(t *testing.T) {
|
||||
provider := &mockLLMProvider{
|
||||
embeddingResponse: embedding,
|
||||
chatResponse: &llm.ChatResponse{Choices: []llm.ChatChoice{{Message: llm.ChatMessage{Content: "General answer."}}}},
|
||||
}
|
||||
repo := &mockResponseRepo{searchByEmbeddingResult: []model.CaptainAssistantResponse{
|
||||
{Question: "Missing", Answer: "Must not ground."},
|
||||
{Question: "Relative", Answer: "Must not ground.", SourceURL: "/knowledge"},
|
||||
}}
|
||||
svc := &CaptainAssistantService{responseRepo: repo, llmProvider: provider, promptBuilder: NewSystemPromptBuilder()}
|
||||
withCitations := &model.CaptainAssistant{Name: "Fin", Status: model.AssistantStatusActive, Config: []byte(`{"feature_citation":true}`)}
|
||||
|
||||
content, err := svc.generatePlaygroundLLMResponse(context.Background(), withCitations, history)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "General answer.", content)
|
||||
assert.NotContains(t, provider.lastChatRequest.Messages[0].Content, "Knowledge Base Context")
|
||||
assert.NotContains(t, provider.lastChatRequest.Messages[0].Content, "[[n](URL)]")
|
||||
})
|
||||
|
||||
t.Run("citation disabled keeps existing Markdown behavior", func(t *testing.T) {
|
||||
provider := &mockLLMProvider{
|
||||
embeddingResponse: embedding,
|
||||
chatResponse: &llm.ChatResponse{Choices: []llm.ChatChoice{{Message: llm.ChatMessage{Content: "See [the guide](/knowledge)."}}}},
|
||||
}
|
||||
repo := &mockResponseRepo{searchByEmbeddingResult: []model.CaptainAssistantResponse{{
|
||||
Question: "Legacy", Answer: "Existing unsourced answer.",
|
||||
}}}
|
||||
svc := &CaptainAssistantService{responseRepo: repo, llmProvider: provider, promptBuilder: NewSystemPromptBuilder()}
|
||||
withoutCitations := &model.CaptainAssistant{Name: "Fin", Status: model.AssistantStatusActive, Config: []byte(`{"feature_citation":false}`)}
|
||||
|
||||
content, err := svc.generatePlaygroundLLMResponse(context.Background(), withoutCitations, history)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "See [the guide](/knowledge).", content)
|
||||
assert.Contains(t, provider.lastChatRequest.Messages[0].Content, "Q: Legacy")
|
||||
assert.NotContains(t, provider.lastChatRequest.Messages[0].Content, "[[n](URL)]")
|
||||
})
|
||||
|
||||
t.Run("embedding provider failure stops chat", func(t *testing.T) {
|
||||
provider := &mockLLMProvider{embeddingError: &net.OpError{Op: "dial", Net: "tcp", Err: errors.New("connection refused")}}
|
||||
svc := &CaptainAssistantService{responseRepo: &mockResponseRepo{}, llmProvider: provider, promptBuilder: NewSystemPromptBuilder()}
|
||||
|
||||
@@ -6,6 +6,7 @@ import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"math"
|
||||
"net/url"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
@@ -818,13 +819,16 @@ func (s *CaptainAssistantService) generatePlaygroundLLMResponse(ctx context.Cont
|
||||
// RAG: embed the latest user message and search approved FAQ responses.
|
||||
// This mirrors Chatwoot's Captain playground which injects knowledge base context.
|
||||
systemPrompt := s.promptBuilder.BuildAssistantPrompt(assistant, cfg)
|
||||
ragContext, err := s.retrieveFAQContext(ctx, assistant.ID, cfg, history)
|
||||
ragContext, err := s.retrieveFAQContext(ctx, assistant.ID, history, cfg.FeatureCitation)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if ragContext != "" {
|
||||
systemPrompt += "\n\n" + ragContext
|
||||
systemPrompt += "\n\nUse the above FAQ entries as reference when answering and cite used entries as [FAQ n]. If the FAQ entries are not relevant, rely on your general knowledge."
|
||||
systemPrompt += "\n\nUse the above FAQ entries as reference when answering. If the FAQ entries are not relevant, rely on your general knowledge."
|
||||
if cfg.FeatureCitation {
|
||||
systemPrompt += " Cite each used sourced FAQ as [[n](URL)] with its exact Source URL; never invent a URL."
|
||||
}
|
||||
}
|
||||
|
||||
messages := []llm.ChatMessage{{Role: "system", Content: systemPrompt}}
|
||||
@@ -869,12 +873,10 @@ func (s *CaptainAssistantService) generatePlaygroundLLMResponse(ctx context.Cont
|
||||
|
||||
// retrieveFAQContext generates an embedding for the latest user message,
|
||||
// searches approved FAQ responses via pgvector, and returns formatted context.
|
||||
// Returns empty string when FAQ is disabled, there is no user query, or there is no matching FAQ.
|
||||
func (s *CaptainAssistantService) retrieveFAQContext(ctx context.Context, assistantID uint, cfg *model.AssistantConfig, history []PlaygroundMessage) (string, error) {
|
||||
if cfg != nil && !cfg.FeatureFAQ {
|
||||
return "", nil
|
||||
}
|
||||
|
||||
// Returns empty string when there is no user query or no matching approved knowledge.
|
||||
// feature_faq controls FAQ generation from resolved conversations upstream; it
|
||||
// does not disable retrieval of already-approved knowledge.
|
||||
func (s *CaptainAssistantService) retrieveFAQContext(ctx context.Context, assistantID uint, history []PlaygroundMessage, requireSource bool) (string, error) {
|
||||
// Extract the latest user message
|
||||
userMsg := ""
|
||||
for i := len(history) - 1; i >= 0; i-- {
|
||||
@@ -886,7 +888,6 @@ func (s *CaptainAssistantService) retrieveFAQContext(ctx context.Context, assist
|
||||
if userMsg == "" {
|
||||
return "", nil
|
||||
}
|
||||
|
||||
// Generate embedding for the question
|
||||
embedResp, err := s.llmProvider.CreateEmbedding(ctx, llm.EmbeddingRequest{
|
||||
Input: []string{userMsg},
|
||||
@@ -917,12 +918,27 @@ func (s *CaptainAssistantService) retrieveFAQContext(ctx context.Context, assist
|
||||
}
|
||||
|
||||
var contextParts []string
|
||||
for i, r := range results {
|
||||
contextParts = append(contextParts, fmt.Sprintf("[FAQ %d]\nQ: %s\nA: %s", i+1, r.Question, r.Answer))
|
||||
for _, r := range results {
|
||||
if requireSource && !validCitationSource(r.SourceURL) {
|
||||
continue
|
||||
}
|
||||
entry := fmt.Sprintf("[FAQ %d]\nQ: %s\nA: %s", len(contextParts)+1, r.Question, r.Answer)
|
||||
if r.SourceURL != "" {
|
||||
entry += "\nSource: " + r.SourceURL
|
||||
}
|
||||
contextParts = append(contextParts, entry)
|
||||
}
|
||||
if len(contextParts) == 0 {
|
||||
return "", nil
|
||||
}
|
||||
return "Knowledge Base Context:\n" + strings.Join(contextParts, "\n\n"), nil
|
||||
}
|
||||
|
||||
func validCitationSource(raw string) bool {
|
||||
u, err := url.ParseRequestURI(strings.TrimSpace(raw))
|
||||
return err == nil && (u.Scheme == "http" || u.Scheme == "https") && u.Host != ""
|
||||
}
|
||||
|
||||
func withAssistantGenerationConfig(ctx context.Context, cfg *model.AssistantConfig) context.Context {
|
||||
if cfg == nil || !cfg.TemperatureConfigured {
|
||||
return ctx
|
||||
|
||||
@@ -39,7 +39,7 @@ func (p *scriptedCaptainSkillProvider) ChatCompletion(_ context.Context, req llm
|
||||
}
|
||||
|
||||
func (*scriptedCaptainSkillProvider) CreateEmbedding(context.Context, llm.EmbeddingRequest) (*llm.EmbeddingResponse, error) {
|
||||
return nil, nil
|
||||
return &llm.EmbeddingResponse{Data: []llm.EmbeddingData{{Embedding: []float64{0.1, 0.2, 0.3}}}}, nil
|
||||
}
|
||||
|
||||
func (*scriptedCaptainSkillProvider) ChatCompletionStream(context.Context, llm.ChatRequest, func(llm.StreamChunk) error) error {
|
||||
@@ -78,7 +78,7 @@ func newCaptainSkillRuntimeDB(t *testing.T, account *model.Account, assistant *m
|
||||
t.Helper()
|
||||
db, err := gorm.Open(sqlite.Open("file:"+strings.ReplaceAll(t.Name(), "/", "_")+"?mode=memory&cache=shared"), &gorm.Config{})
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, db.AutoMigrate(&model.Account{}, &model.CaptainAssistant{}, &model.CaptainCustomTool{}, &model.CaptainSkill{}, &model.CaptainSkillReference{}, &model.CaptainAssistantSkill{}))
|
||||
require.NoError(t, db.AutoMigrate(&model.Account{}, &model.CaptainAssistant{}, &model.CaptainDocument{}, &model.CaptainAssistantResponse{}, &model.CaptainCustomTool{}, &model.CaptainSkill{}, &model.CaptainSkillReference{}, &model.CaptainAssistantSkill{}))
|
||||
require.NoError(t, db.Create(account).Error)
|
||||
assistant.ID = 0
|
||||
require.NoError(t, db.Create(assistant).Error)
|
||||
|
||||
@@ -179,7 +179,7 @@ func fixtureChatMessage(req chatRequest) (chatMessage, string) {
|
||||
}
|
||||
}
|
||||
if strings.Contains(all, "Knowledge Base Context:") {
|
||||
return chatMessage{Role: "assistant", Content: "The local fixture warranty window is 37 days [FAQ 1]."}, "stop"
|
||||
return chatMessage{Role: "assistant", Content: "The local fixture warranty window is 37 days [[1](http://captain-fixture:8080/knowledge)]."}, "stop"
|
||||
}
|
||||
return chatMessage{Role: "assistant", Content: "Fixture chat response."}, "stop"
|
||||
}
|
||||
|
||||
@@ -52,6 +52,10 @@ func TestAcceptanceFixtureContract(t *testing.T) {
|
||||
if status := post(t, server.URL+"/v1/chat/completions", faq, &response); status != http.StatusOK || len(response.Choices) != 1 || !strings.Contains(response.Choices[0].Message.Content, `"faqs"`) {
|
||||
t.Fatalf("faq status=%d choices=%d", status, len(response.Choices))
|
||||
}
|
||||
grounded := `{"model":"gpt-5.6-luna","messages":[{"role":"system","content":"Knowledge Base Context:\nSource: http://captain-fixture:8080/knowledge"},{"role":"user","content":"warranty?"}]}`
|
||||
if status := post(t, server.URL+"/v1/chat/completions", grounded, &response); status != http.StatusOK || !strings.Contains(response.Choices[0].Message.Content, "[[1](http://captain-fixture:8080/knowledge)]") {
|
||||
t.Fatalf("grounded status=%d content=%q", status, response.Choices[0].Message.Content)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("skill activation", func(t *testing.T) {
|
||||
|
||||
Reference in New Issue
Block a user