package runtime import ( "context" "strings" "sync/atomic" "testing" "time" "code.tczkiot.com/wlw/ai-agent/contract" "code.tczkiot.com/wlw/ai-agent/internal/models" "code.tczkiot.com/wlw/ai-agent/internal/pkg/enums" "code.tczkiot.com/wlw/ai-agent/internal/repositories" "github.com/glebarez/sqlite" "github.com/mlogclub/simple/sqls" "gorm.io/gorm" "gorm.io/gorm/schema" ) func TestTriggerReplyAsyncBindsCustomerProofToCurrentMessage(t *testing.T) { database, err := gorm.Open(sqlite.Open("file:"+strings.ReplaceAll(t.Name(), "/", "_")+"?mode=memory&cache=shared"), &gorm.Config{ NamingStrategy: schema.NamingStrategy{TablePrefix: "t_", SingularTable: true}, }) if err != nil { t.Fatalf("open sqlite: %v", err) } if err := database.AutoMigrate(&models.AIAgent{}, &models.AgentToolInvocation{}); err != nil { t.Fatalf("migrate runtime claim tables: %v", err) } sqls.SetDB(database) agent := models.AIAgent{Name: "test", Status: enums.StatusOk, PublishedRevisionID: 19, ReplyTimeoutSeconds: 5} if err := database.Create(&agent).Error; err != nil { t.Fatalf("create agent: %v", err) } proofContext := contract.WithCustomerAccessProof(context.Background(), contract.CustomerAccessProof{ SessionID: "opaque-session", TargetType: "device", TargetID: 27, ExpiresAt: time.Now().Add(15 * time.Minute), }) conversation := models.Conversation{ID: 101, AIAgentID: agent.ID} message := models.Message{ ID: 202, ConversationID: conversation.ID, SenderType: enums.IMSenderTypeCustomer, Content: "请帮我切换网络", RequestID: "request-303", } received := make(chan contract.CustomerAccessProof, 1) service := newAIReplyService() service.triggerReply = func(ctx context.Context, _ models.Conversation, _ models.Message, _ models.AIAgent) error { proof, ok := contract.CustomerAccessProofFromContext(ctx) if !ok { return context.Canceled } received <- proof return nil } service.TriggerReplyAsync(proofContext, conversation, message) select { case proof := <-received: if proof.ConversationID != conversation.ID || proof.MessageID != message.ID || proof.RequestID != message.RequestID { t.Fatalf("async proof was not bound to current message: %#v", proof) } case <-time.After(time.Second): t.Fatal("reply execution did not receive customer access proof") } } func TestTriggerReplyAsyncClaimsMessageRevisionOnce(t *testing.T) { database, err := gorm.Open(sqlite.Open("file:"+strings.ReplaceAll(t.Name(), "/", "_")+"?mode=memory&cache=shared"), &gorm.Config{ NamingStrategy: schema.NamingStrategy{TablePrefix: "t_", SingularTable: true}, }) if err != nil { t.Fatalf("open sqlite: %v", err) } if err := database.AutoMigrate(&models.AIAgent{}, &models.AgentToolInvocation{}); err != nil { t.Fatalf("migrate runtime claim tables: %v", err) } sqls.SetDB(database) agent := models.AIAgent{Name: "test", Status: enums.StatusOk, PublishedRevisionID: 7, ReplyTimeoutSeconds: 5} if err := database.Create(&agent).Error; err != nil { t.Fatalf("create agent: %v", err) } var executions atomic.Int32 started := make(chan struct{}) release := make(chan struct{}) done := make(chan struct{}) service := newAIReplyService() service.triggerReply = func(context.Context, models.Conversation, models.Message, models.AIAgent) error { if executions.Add(1) == 1 { close(started) } <-release close(done) return nil } conversation := models.Conversation{ID: 100, AIAgentID: agent.ID} message := models.Message{ID: 200, ConversationID: conversation.ID, SenderType: enums.IMSenderTypeCustomer, Content: "hello", RequestID: "req-concurrent"} service.TriggerReplyAsync(context.Background(), conversation, message) service.TriggerReplyAsync(context.Background(), conversation, message) select { case <-started: case <-time.After(time.Second): t.Fatal("reply execution did not start") } if got := executions.Load(); got != 1 { t.Fatalf("concurrent triggers executed %d times", got) } close(release) select { case <-done: case <-time.After(time.Second): t.Fatal("reply execution did not finish") } deadline := time.Now().Add(time.Second) for { item := repositories.AgentToolInvocationRepository.GetByIdempotencyKey(database, conversation.ID, aiReplyInvocationToolCode, "message:200:revision:7") if item != nil && item.Status == "completed" { break } if time.Now().After(deadline) { t.Fatalf("reply invocation was not completed: %#v", item) } time.Sleep(5 * time.Millisecond) } } func TestTriggerReplyAsyncReconcilesRecoveredCommittedReplyWithoutModelCall(t *testing.T) { database, err := gorm.Open(sqlite.Open("file:"+strings.ReplaceAll(t.Name(), "/", "_")+"?mode=memory&cache=shared"), &gorm.Config{ NamingStrategy: schema.NamingStrategy{TablePrefix: "t_", SingularTable: true}, }) if err != nil { t.Fatalf("open sqlite: %v", err) } if err := database.AutoMigrate(&models.AIAgent{}, &models.AgentToolInvocation{}, &models.Message{}); err != nil { t.Fatalf("migrate runtime claim tables: %v", err) } sqls.SetDB(database) agent := models.AIAgent{Name: "test", Status: enums.StatusOk, PublishedRevisionID: 8, ReplyTimeoutSeconds: 1} if err := database.Create(&agent).Error; err != nil { t.Fatalf("create agent: %v", err) } conversation := models.Conversation{ID: 101, AIAgentID: agent.ID} message := models.Message{ID: 201, ConversationID: conversation.ID, SenderType: enums.IMSenderTypeCustomer, Content: "hello"} invocation := models.AgentToolInvocation{ ConversationID: conversation.ID, AIAgentID: agent.ID, ToolCode: aiReplyInvocationToolCode, IdempotencyKey: "message:201:revision:8", Status: "running", ResultData: "old-lease", } if err := database.Create(&invocation).Error; err != nil { t.Fatalf("create stale invocation: %v", err) } if err := database.Model(&models.AgentToolInvocation{}).Where("id = ?", invocation.ID).Update("updated_at", time.Now().Add(-time.Hour)).Error; err != nil { t.Fatalf("age invocation: %v", err) } committed := models.Message{ConversationID: conversation.ID, ClientMsgID: "ai_reply_201", SenderType: enums.IMSenderTypeAI, MessageType: enums.IMMessageTypeText, Content: "done"} if err := database.Create(&committed).Error; err != nil { t.Fatalf("create committed reply: %v", err) } var executions atomic.Int32 service := newAIReplyService() service.triggerReply = func(context.Context, models.Conversation, models.Message, models.AIAgent) error { executions.Add(1) return nil } service.TriggerReplyAsync(context.Background(), conversation, message) if executions.Load() != 0 { t.Fatalf("model executed despite committed reply: %d", executions.Load()) } item := repositories.AgentToolInvocationRepository.GetByIdempotencyKey(database, conversation.ID, aiReplyInvocationToolCode, invocation.IdempotencyKey) if item == nil || item.Status != "completed" { t.Fatalf("recovered invocation not reconciled: %#v", item) } }