Files
ai-agent/internal/ai/runtime/reply_commit_service_test.go
T

179 lines
5.5 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package runtime
import (
"context"
"strings"
"testing"
"time"
"code.tczkiot.com/wlw/ai-agent/internal/models"
"code.tczkiot.com/wlw/ai-agent/internal/pkg/enums"
"github.com/glebarez/sqlite"
"github.com/mlogclub/simple/sqls"
"gorm.io/gorm"
"gorm.io/gorm/schema"
)
func TestReplyCommitStoresAIMessage(t *testing.T) {
db := setupReplyCommitTestDB(t)
aiAgent := createReplyCommitTestAIAgent(t, db)
conversation := createReplyCommitTestConversation(t, db, aiAgent.ID)
replyMessage, err := newReplyCommitService().CommitAIReply(replyCommitInput{
Conversation: *conversation,
Message: models.Message{ID: 101, RequestID: "trace-101"},
AIAgent: *aiAgent,
ReplyText: "AI reply",
ClientPrefix: "ai_reply",
})
if err != nil {
t.Fatalf("CommitAIReply() error = %v", err)
}
if replyMessage == nil {
t.Fatalf("expected reply message")
}
var stored models.Message
if err := db.First(&stored, replyMessage.ID).Error; err != nil {
t.Fatalf("find reply message: %v", err)
}
if stored.Content != "AI reply" || stored.RequestID != "trace-101" {
t.Fatalf("unexpected stored reply: %#v", stored)
}
}
func TestReplyCommitRejectsSensitiveModelOutput(t *testing.T) {
db := setupReplyCommitTestDB(t)
aiAgent := createReplyCommitTestAIAgent(t, db)
conversation := createReplyCommitTestConversation(t, db, aiAgent.ID)
_, err := newReplyCommitService().CommitAIReply(replyCommitInput{
Conversation: *conversation, Message: models.Message{ID: 102, RequestID: "trace-102"}, AIAgent: *aiAgent,
ReplyText: "authorization=Bearer-secret", ClientPrefix: "ai_reply",
})
if err == nil {
t.Fatal("expected sensitive model output to be rejected")
}
var count int64
if err := db.Model(&models.Message{}).Where("conversation_id = ?", conversation.ID).Count(&count).Error; err != nil {
t.Fatalf("count messages: %v", err)
}
if count != 0 {
t.Fatalf("unexpected message written for rejected output: %d", count)
}
}
func TestReplyCommitRedactsICCIDBeforePersisting(t *testing.T) {
db := setupReplyCommitTestDB(t)
aiAgent := createReplyCommitTestAIAgent(t, db)
conversation := createReplyCommitTestConversation(t, db, aiAgent.ID)
replyMessage, err := newReplyCommitService().CommitAIReply(replyCommitInput{
Conversation: *conversation, Message: models.Message{ID: 104, RequestID: "trace-104"}, AIAgent: *aiAgent,
ReplyText: "业务状态:正常\nICCID8986042302268012345", ClientPrefix: "ai_reply",
})
if err != nil {
t.Fatalf("CommitAIReply() error = %v", err)
}
var stored models.Message
if err := db.First(&stored, replyMessage.ID).Error; err != nil {
t.Fatalf("find reply message: %v", err)
}
if strings.Contains(stored.Content, "8986042302268012345") || !strings.Contains(stored.Content, "系统内部标识") {
t.Fatalf("ICCID was persisted in customer reply: %q", stored.Content)
}
}
func TestFailureReplyDeduplicatesByDeterministicClientMessageID(t *testing.T) {
db := setupReplyCommitTestDB(t)
aiAgent := createReplyCommitTestAIAgent(t, db)
conversation := createReplyCommitTestConversation(t, db, aiAgent.ID)
message := models.Message{ID: 103, ConversationID: conversation.ID, RequestID: "trace-shared"}
service := newAIReplyService()
service.commitFailureReplyIfNeeded(*conversation, message, *aiAgent, context.DeadlineExceeded)
// The request ID is deliberately changed: error idempotency is tied to the
// triggering customer message, not a transport trace that can be regenerated.
message.RequestID = "trace-retry"
service.commitFailureReplyIfNeeded(*conversation, message, *aiAgent, context.DeadlineExceeded)
var messages []models.Message
if err := db.Where("conversation_id = ? AND client_msg_id = ?", conversation.ID, "ai_error_103").Find(&messages).Error; err != nil {
t.Fatalf("find failure messages: %v", err)
}
if len(messages) != 1 {
t.Fatalf("failure reply count = %d, want 1", len(messages))
}
}
func setupReplyCommitTestDB(t *testing.T) *gorm.DB {
t.Helper()
dbName := "reply_commit_test_" + strings.NewReplacer("/", "_").Replace(t.Name())
db, err := gorm.Open(sqlite.Open("file:"+dbName+"?mode=memory&cache=shared"), &gorm.Config{
NamingStrategy: schema.NamingStrategy{
TablePrefix: "t_",
SingularTable: true,
},
})
if err != nil {
t.Fatalf("open sqlite db: %v", err)
}
sqlDB, err := db.DB()
if err != nil {
t.Fatalf("get sqlite db: %v", err)
}
t.Cleanup(func() {
if err := sqlDB.Close(); err != nil {
t.Fatalf("close sqlite db: %v", err)
}
})
if err := db.AutoMigrate(
&models.AIAgent{},
&models.Channel{},
&models.ChannelMessageOutbox{},
&models.Conversation{},
&models.ConversationReadState{},
&models.ConversationEventLog{},
&models.Message{},
); err != nil {
t.Fatalf("auto migrate: %v", err)
}
sqls.SetDB(db)
return db
}
func createReplyCommitTestAIAgent(t *testing.T, db *gorm.DB) *models.AIAgent {
t.Helper()
now := time.Now()
item := &models.AIAgent{
Name: "reply-agent",
Status: enums.StatusOk,
AuditFields: models.AuditFields{
CreatedAt: now,
UpdatedAt: now,
},
}
if err := db.Create(item).Error; err != nil {
t.Fatalf("create ai agent: %v", err)
}
return item
}
func createReplyCommitTestConversation(t *testing.T, db *gorm.DB, aiAgentID int64) *models.Conversation {
t.Helper()
now := time.Now()
item := &models.Conversation{
CustomerID: 1,
ChannelID: 11,
AIAgentID: aiAgentID,
Status: enums.IMConversationStatusAIServing,
LastActiveAt: now,
AuditFields: models.AuditFields{
CreatedAt: now,
UpdatedAt: now,
},
}
if err := db.Create(item).Error; err != nil {
t.Fatalf("create conversation: %v", err)
}
return item
}