diff --git a/internal/ai/runtime/executor/answerability_gate_test.go b/internal/ai/runtime/executor/answerability_gate_test.go index 4520a9a..d795a60 100644 --- a/internal/ai/runtime/executor/answerability_gate_test.go +++ b/internal/ai/runtime/executor/answerability_gate_test.go @@ -186,6 +186,46 @@ func TestKnowledgeAnswerabilityGateEvaluateAllowsAnswerableDecisionAndProducesKn } } +func TestBuildRunMessagesReturnsFallbackWhenGateRejectsGrayZoneQuestion(t *testing.T) { + summary := &RunResult{} + question := "满足什么条件可以退款?" + gate := newTestKnowledgeAnswerabilityGate(newAnswerabilityRetrieverWithHit(), &fakeAnswerabilityChatModel{ + response: `{"answerable": false, "reason": "retrieved snippets mention refunds but not the requested condition", "missingInfo": ["refund condition"]}`, + }) + + messages := buildRunMessages(context.Background(), newAnswerabilityGateRunInput(question, "1"), summary, nil, gate) + + if !strings.Contains(summary.ReplyText, "建议你联系人工客服进一步确认。") { + t.Fatalf("expected human-support fallback, got %q", summary.ReplyText) + } + if messagesContainContent(messages, question) { + t.Fatalf("expected returned messages to omit current user message, got %#v", messages) + } +} + +func TestBuildRunMessagesInjectsKnowledgeWhenGateAllows(t *testing.T) { + summary := &RunResult{} + question := "满足什么条件可以退款?" + gate := newTestKnowledgeAnswerabilityGate(newAnswerabilityRetrieverWithHit(), &fakeAnswerabilityChatModel{ + response: `{"answerable": true, "reason": "refund condition is directly supported", "supportingChunkIds": ["101"]}`, + }) + + messages := buildRunMessages(context.Background(), newAnswerabilityGateRunInput(question, "1"), summary, nil, gate) + + if summary.ReplyText != "" { + t.Fatalf("expected no fallback, got %q", summary.ReplyText) + } + if !messagesContainContent(messages, "知识库回答约束") { + t.Fatalf("expected knowledge instruction in messages, got %#v", messages) + } + if !messagesContainContent(messages, "购买后七天内且未使用可以退款。") { + t.Fatalf("expected knowledge context in messages, got %#v", messages) + } + if !messagesContainContent(messages, question) { + t.Fatalf("expected current user message in messages, got %#v", messages) + } +} + func newTestKnowledgeAnswerabilityGate(retriever knowledgeContextRetriever, chatModel model.BaseChatModel) *KnowledgeAnswerabilityGate { return &KnowledgeAnswerabilityGate{ newRetriever: func(aiAgent models.AIAgent) knowledgeContextRetriever { @@ -257,3 +297,12 @@ func (f *fakeAnswerabilityChatModel) Generate(ctx context.Context, input []*sche func (f *fakeAnswerabilityChatModel) Stream(ctx context.Context, input []*schema.Message, opts ...model.Option) (*schema.StreamReader[*schema.Message], error) { return nil, errors.New("stream is not implemented in fakeAnswerabilityChatModel") } + +func messagesContainContent(messages []*schema.Message, text string) bool { + for _, message := range messages { + if message != nil && strings.Contains(message.Content, text) { + return true + } + } + return false +} diff --git a/internal/ai/runtime/executor/context_builders.go b/internal/ai/runtime/executor/context_builders.go index bf0ca08..fc314e0 100644 --- a/internal/ai/runtime/executor/context_builders.go +++ b/internal/ai/runtime/executor/context_builders.go @@ -6,13 +6,12 @@ import ( "cs-agent/internal/ai/runtime/internal/impl/adapter" "cs-agent/internal/ai/runtime/internal/impl/callbacks" - "cs-agent/internal/ai/runtime/internal/impl/retrievers" "cs-agent/internal/pkg/utils" "github.com/cloudwego/eino/schema" ) -func buildRunMessages(ctx context.Context, req RunInput, summary *RunResult, collector *callbacks.RuntimeTraceCollector) []*schema.Message { +func buildRunMessages(ctx context.Context, req RunInput, summary *RunResult, collector *callbacks.RuntimeTraceCollector, gate *KnowledgeAnswerabilityGate) []*schema.Message { history := adapter.BuildHistoryMessages(req.Conversation.ID, req.UserMessage.ID, 12) if summary != nil { summary.HistoryMessageCount = len(history.Messages) @@ -24,7 +23,7 @@ func buildRunMessages(ctx context.Context, req RunInput, summary *RunResult, col } messages := make([]*schema.Message, 0, len(history.Messages)+3) messages = append(messages, history.Messages...) - decision := appendRetrievedContext(ctx, req, summary, collector, &messages) + decision := appendRetrievedContext(ctx, req, summary, collector, gate, &messages) if strings.TrimSpace(decision.FallbackReply) != "" { if summary != nil { summary.ReplyText = decision.FallbackReply @@ -35,30 +34,32 @@ func buildRunMessages(ctx context.Context, req RunInput, summary *RunResult, col return messages } -func appendRetrievedContext(ctx context.Context, req RunInput, summary *RunResult, collector *callbacks.RuntimeTraceCollector, messages *[]*schema.Message) knowledgeGuardDecision { +func appendRetrievedContext(ctx context.Context, req RunInput, summary *RunResult, collector *callbacks.RuntimeTraceCollector, gate *KnowledgeAnswerabilityGate, messages *[]*schema.Message) knowledgeGuardDecision { if messages == nil { return knowledgeGuardDecision{} } - retriever := retrievers.NewKnowledgeRetriever(req.AIAgent) - retrieveOptions := retrievers.DefaultKnowledgeRetrieveOptions() - retrieveOptions.QueryPreview = preview(req.UserMessage.Content, 120) - retrieveResult, retrieveErr := retriever.RetrieveContextByOptions(ctx, retrieveOptions, strings.TrimSpace(req.UserMessage.Content)) - if retrieveErr != nil || retrieveResult == nil { - return buildKnowledgeUnavailableDecision(req.AIAgent, retriever.KnowledgeBaseIDs()) + if gate == nil { + gate = NewKnowledgeAnswerabilityGate() } - if summary != nil { - summary.RetrieverCount = len(retrieveResult.Hits) + state, err := gate.Evaluate(ctx, answerabilityGateInput{ + Request: req, + Summary: summary, + Collector: collector, + Messages: append([]*schema.Message(nil), (*messages)...), + }) + if err != nil || state == nil { + decision := buildKnowledgeUnavailableDecision(req.AIAgent, utils.SplitInt64s(req.AIAgent.KnowledgeIDs)) + if strings.TrimSpace(decision.FallbackReply) != "" { + decision.FallbackReply = resolveKnowledgeHumanSupportFallback(req.AIAgent) + } + return decision } - if collector != nil { - collector.SetRetrieverSummary(retrieveResult.TraceSummary) - collector.AddRetrieverItems(retrieveResult.TraceItems) + if strings.TrimSpace(state.FallbackReply) != "" { + return knowledgeGuardDecision{FallbackReply: state.FallbackReply} } - decision := buildKnowledgeGuardDecision(req.AIAgent, retrieveResult) - if len(decision.Instructions) > 0 { - *messages = append(*messages, decision.Instructions...) + if state.SkipGate { + return knowledgeGuardDecision{} } - if strings.TrimSpace(retrieveResult.ContextText) != "" { - *messages = append(*messages, schema.SystemMessage(retrieveResult.ContextText)) - } - return decision + *messages = append((*messages)[:0], state.Input.Messages...) + return state.Decision } diff --git a/internal/ai/runtime/executor/service.go b/internal/ai/runtime/executor/service.go index 79cb13e..9129724 100644 --- a/internal/ai/runtime/executor/service.go +++ b/internal/ai/runtime/executor/service.go @@ -13,14 +13,16 @@ import ( ) type Service struct { - agentFactory *factory.AgentFactory - runnerFactory *factory.RunnerFactory + agentFactory *factory.AgentFactory + runnerFactory *factory.RunnerFactory + answerabilityGate *KnowledgeAnswerabilityGate } func NewService() *Service { return &Service{ - agentFactory: factory.NewAgentFactory(), - runnerFactory: factory.NewRunnerFactory(), + agentFactory: factory.NewAgentFactory(), + runnerFactory: factory.NewRunnerFactory(), + answerabilityGate: NewKnowledgeAnswerabilityGate(), } } @@ -84,7 +86,7 @@ func (s *Service) ExecuteRun(ctx context.Context, req RunInput) (*RunResult, err summary.TraceData = collector.Marshal() return summary, fmt.Errorf("%s", summary.ErrorMessage) } - messages := buildRunMessages(ctx, req, summary, collector) + messages := buildRunMessages(ctx, req, summary, collector, s.answerabilityGate) if strings.TrimSpace(summary.ReplyText) != "" { summary.Status = "completed" summary.ModelName = req.AIConfig.ModelName