diff --git a/internal/ai/runtime/executor/answerability_gate.go b/internal/ai/runtime/executor/answerability_gate.go index 4b70936..1ef82fd 100644 --- a/internal/ai/runtime/executor/answerability_gate.go +++ b/internal/ai/runtime/executor/answerability_gate.go @@ -2,6 +2,9 @@ package executor import ( "context" + "encoding/json" + "fmt" + "strings" "time" "cs-agent/internal/ai/runtime/internal/impl/callbacks" @@ -10,6 +13,8 @@ import ( "cs-agent/internal/models" "github.com/cloudwego/eino/components/model" + "github.com/cloudwego/eino/components/prompt" + "github.com/cloudwego/eino/compose" "github.com/cloudwego/eino/schema" ) @@ -76,3 +81,341 @@ func NewKnowledgeAnswerabilityGate() *KnowledgeAnswerabilityGate { now: time.Now, } } + +func (g *KnowledgeAnswerabilityGate) withDefaults() *KnowledgeAnswerabilityGate { + if g == nil { + return NewKnowledgeAnswerabilityGate() + } + ret := *g + defaults := NewKnowledgeAnswerabilityGate() + if ret.newRetriever == nil { + ret.newRetriever = defaults.newRetriever + } + if ret.newChatModel == nil { + ret.newChatModel = defaults.newChatModel + } + if ret.now == nil { + ret.now = time.Now + } + return &ret +} + +func (g *KnowledgeAnswerabilityGate) Evaluate(ctx context.Context, input answerabilityGateInput) (*answerabilityGateState, error) { + gate := g.withDefaults() + graph := compose.NewGraph[*answerabilityGateState, *answerabilityGateState]() + if err := graph.AddLambdaNode(answerabilityNodeRetrieve, compose.InvokableLambda(gate.retrieveKnowledge)); err != nil { + return nil, err + } + if err := graph.AddLambdaNode(answerabilityNodeGrade, compose.InvokableLambda(gate.gradeAnswerability)); err != nil { + return nil, err + } + if err := graph.AddLambdaNode(answerabilityNodeAllow, compose.InvokableLambda(allowAnswerabilityPassThrough)); err != nil { + return nil, err + } + if err := graph.AddLambdaNode(answerabilityNodeFallback, compose.InvokableLambda(fallbackAnswerabilityPassThrough)); err != nil { + return nil, err + } + if err := graph.AddEdge(compose.START, answerabilityNodeRetrieve); err != nil { + return nil, err + } + if err := graph.AddEdge(answerabilityNodeRetrieve, answerabilityNodeGrade); err != nil { + return nil, err + } + if err := graph.AddBranch(answerabilityNodeGrade, compose.NewGraphBranch(routeAnswerabilityGate, map[string]bool{ + answerabilityNodeAllow: true, + answerabilityNodeFallback: true, + })); err != nil { + return nil, err + } + if err := graph.AddEdge(answerabilityNodeAllow, compose.END); err != nil { + return nil, err + } + if err := graph.AddEdge(answerabilityNodeFallback, compose.END); err != nil { + return nil, err + } + runnable, err := graph.Compile(ctx) + if err != nil { + return nil, err + } + return runnable.Invoke(ctx, &answerabilityGateState{Input: input}) +} + +func routeAnswerabilityGate(ctx context.Context, state *answerabilityGateState) (string, error) { + if state == nil { + return answerabilityNodeFallback, nil + } + if state.SkipGate || strings.TrimSpace(state.FallbackReply) == "" { + return answerabilityNodeAllow, nil + } + return answerabilityNodeFallback, nil +} + +func allowAnswerabilityPassThrough(ctx context.Context, state *answerabilityGateState) (*answerabilityGateState, error) { + if state == nil { + return &answerabilityGateState{}, nil + } + if len(state.Decision.Instructions) > 0 { + state.Input.Messages = append(state.Input.Messages, state.Decision.Instructions...) + } + if state.RetrieveResult != nil { + if contextText := strings.TrimSpace(state.RetrieveResult.ContextText); contextText != "" { + state.Input.Messages = append(state.Input.Messages, schema.SystemMessage(contextText)) + } + } + return state, nil +} + +func fallbackAnswerabilityPassThrough(ctx context.Context, state *answerabilityGateState) (*answerabilityGateState, error) { + if state == nil { + return &answerabilityGateState{}, nil + } + return state, nil +} + +func (g *KnowledgeAnswerabilityGate) retrieveKnowledge(ctx context.Context, state *answerabilityGateState) (*answerabilityGateState, error) { + if state == nil { + state = &answerabilityGateState{} + } + gate := g.withDefaults() + req := state.Input.Request + retriever := gate.newRetriever(req.AIAgent) + if retriever == nil { + state.SkipGate = true + state.recordAnswerability(answerabilityStatusSkipped, "knowledge retriever unavailable", nil) + return state, nil + } + knowledgeIDs := retriever.KnowledgeBaseIDs() + state.KnowledgeIDs = append([]int64(nil), knowledgeIDs...) + if len(knowledgeIDs) == 0 { + state.SkipGate = true + state.recordAnswerability(answerabilityStatusSkipped, "no knowledge configured", nil) + return state, nil + } + query := strings.TrimSpace(req.UserMessage.Content) + if query == "" { + state.FallbackReply = resolveKnowledgeHumanSupportFallback(req.AIAgent) + state.recordAnswerability(answerabilityStatusUnanswerable, "empty user question", nil) + return state, nil + } + retrieveOptions := retrievers.DefaultKnowledgeRetrieveOptions() + retrieveOptions.QueryPreview = preview(req.UserMessage.Content, 120) + result, err := retriever.RetrieveContextByOptions(ctx, retrieveOptions, query) + if err != nil { + state.FallbackReply = resolveKnowledgeHumanSupportFallback(req.AIAgent) + state.ErrorMessage = err.Error() + state.recordAnswerability(answerabilityStatusUnanswerable, "knowledge retrieval failed", err) + return state, nil + } + state.RetrieveResult = result + if state.Input.Summary != nil && result != nil { + state.Input.Summary.RetrieverCount = len(result.Hits) + } + if state.Input.Collector != nil && result != nil { + state.Input.Collector.SetRetrieverSummary(result.TraceSummary) + state.Input.Collector.Data.Retriever.Items = append(state.Input.Collector.Data.Retriever.Items, result.TraceItems...) + } + if result == nil || len(result.Hits) == 0 || strings.TrimSpace(result.ContextText) == "" { + state.FallbackReply = resolveKnowledgeHumanSupportFallback(req.AIAgent) + state.recordAnswerability(answerabilityStatusUnanswerable, "no retrieved context", nil) + return state, nil + } + return state, nil +} + +func (g *KnowledgeAnswerabilityGate) gradeAnswerability(ctx context.Context, state *answerabilityGateState) (*answerabilityGateState, error) { + if state == nil { + return &answerabilityGateState{}, nil + } + if state.SkipGate || strings.TrimSpace(state.FallbackReply) != "" { + return state, nil + } + gate := g.withDefaults() + started := time.Now() + req := state.Input.Request + modelInstance, err := gate.newChatModel(ctx, req.AIConfig) + if err != nil { + state.FallbackReply = resolveKnowledgeHumanSupportFallback(req.AIAgent) + state.ErrorMessage = err.Error() + state.recordAnswerabilityWithLatency(answerabilityStatusUnanswerable, "answerability model factory failed", err, started) + return state, nil + } + if modelInstance == nil { + err = fmt.Errorf("answerability model is nil") + state.FallbackReply = resolveKnowledgeHumanSupportFallback(req.AIAgent) + state.ErrorMessage = err.Error() + state.recordAnswerabilityWithLatency(answerabilityStatusUnanswerable, "answerability model factory failed", err, started) + return state, nil + } + messages, err := buildAnswerabilityMessages(ctx, req.UserMessage.Content, state.RetrieveResult) + if err != nil { + state.FallbackReply = resolveKnowledgeHumanSupportFallback(req.AIAgent) + state.ErrorMessage = err.Error() + state.recordAnswerabilityWithLatency(answerabilityStatusUnanswerable, "answerability prompt failed", err, started) + return state, nil + } + response, err := modelInstance.Generate(ctx, messages) + if err != nil { + state.FallbackReply = resolveKnowledgeHumanSupportFallback(req.AIAgent) + state.ErrorMessage = err.Error() + state.recordAnswerabilityWithLatency(answerabilityStatusUnanswerable, "answerability model generate failed", err, started) + return state, nil + } + if response == nil { + err = fmt.Errorf("answerability model returned empty response") + state.FallbackReply = resolveKnowledgeHumanSupportFallback(req.AIAgent) + state.ErrorMessage = err.Error() + state.recordAnswerabilityWithLatency(answerabilityStatusUnanswerable, "answerability model generate failed", err, started) + return state, nil + } + decision, err := parseAnswerabilityDecision(response.Content) + if err != nil { + state.FallbackReply = resolveKnowledgeHumanSupportFallback(req.AIAgent) + state.ErrorMessage = err.Error() + state.recordAnswerabilityWithLatency(answerabilityStatusUnanswerable, "answerability decision parse failed", err, started) + return state, nil + } + state.Grade = decision + if !decision.Answerable { + state.FallbackReply = resolveKnowledgeHumanSupportFallback(req.AIAgent) + state.recordAnswerabilityWithLatency(answerabilityStatusUnanswerable, decision.Reason, nil, started) + return state, nil + } + state.Decision = buildKnowledgeGuardDecision(req.AIAgent, state.RetrieveResult) + if strings.TrimSpace(state.Decision.FallbackReply) != "" { + state.FallbackReply = resolveKnowledgeHumanSupportFallback(req.AIAgent) + state.recordAnswerabilityWithLatency(answerabilityStatusUnanswerable, "knowledge guard rejected retrieved context", nil, started) + return state, nil + } + state.recordAnswerabilityWithLatency(answerabilityStatusAnswerable, decision.Reason, nil, started) + return state, nil +} + +func buildAnswerabilityMessages(ctx context.Context, question string, result *retrievers.KnowledgeRetrieveResult) ([]*schema.Message, error) { + contextText := buildAnswerabilityContext(result) + template := prompt.FromMessages(schema.FString, + schema.SystemMessage(strings.TrimSpace(`你是一个知识库可回答性判定器。 +你只判断“已召回知识片段”是否直接支持回答用户问题,不要使用模型常识补充。 +如果问题中的具体对象、条件、步骤、承诺或限制不能被片段直接支持,判定为不可回答。 +只输出 JSON,不要输出 Markdown、解释或多余文本。 +JSON 字段必须包含: +- answerable: boolean +- reason: string +- supportingChunkIds: string array,answerable 为 true 时必须至少包含一个直接支持的 chunk id +- missingInfo: string array,answerable 为 false 时列出缺失信息`)), + schema.UserMessage(strings.TrimSpace(`用户问题: +{question} + +已召回知识片段: +{context} + +请基于上述片段判定是否可以直接回答用户问题。`)), + ) + return template.Format(ctx, map[string]any{ + "question": strings.TrimSpace(question), + "context": contextText, + }) +} + +func buildAnswerabilityContext(result *retrievers.KnowledgeRetrieveResult) string { + if result == nil { + return "" + } + items := result.ContextResults + if len(items) == 0 { + items = result.Hits + } + if len(items) == 0 { + return strings.TrimSpace(result.ContextText) + } + var builder strings.Builder + for idx, item := range items { + if idx > 0 { + builder.WriteString("\n\n") + } + builder.WriteString(fmt.Sprintf("snippet %d\nknowledgeBaseId: %d\ndocumentId: %d\nchunkId: %d\nscore: %.4f\ncontent:\n%s", + idx+1, + item.KnowledgeBaseID, + item.DocumentID, + item.ChunkID, + item.Score, + strings.TrimSpace(item.Content), + )) + } + return strings.TrimSpace(builder.String()) +} + +func parseAnswerabilityDecision(raw string) (answerabilityDecision, error) { + text := trimMarkdownFence(raw) + if text == "" { + return answerabilityDecision{}, fmt.Errorf("answerability decision is empty") + } + var decision answerabilityDecision + if err := json.Unmarshal([]byte(text), &decision); err != nil { + return answerabilityDecision{}, fmt.Errorf("parse answerability decision: %w", err) + } + decision.Reason = strings.TrimSpace(decision.Reason) + decision.SupportingChunkIDs = trimStringSlice(decision.SupportingChunkIDs) + decision.MissingInfo = trimStringSlice(decision.MissingInfo) + if decision.Answerable && len(decision.SupportingChunkIDs) == 0 { + return answerabilityDecision{}, fmt.Errorf("answerable decision requires supportingChunkIds") + } + return decision, nil +} + +func trimMarkdownFence(raw string) string { + text := strings.TrimSpace(raw) + if !strings.HasPrefix(text, "```") { + return text + } + lines := strings.Split(text, "\n") + if len(lines) == 0 { + return text + } + if strings.HasPrefix(strings.TrimSpace(lines[0]), "```") { + lines = lines[1:] + } + if len(lines) > 0 && strings.HasPrefix(strings.TrimSpace(lines[len(lines)-1]), "```") { + lines = lines[:len(lines)-1] + } + return strings.TrimSpace(strings.Join(lines, "\n")) +} + +func trimStringSlice(items []string) []string { + if len(items) == 0 { + return nil + } + ret := make([]string, 0, len(items)) + for _, item := range items { + item = strings.TrimSpace(item) + if item == "" { + continue + } + ret = append(ret, item) + } + return ret +} + +func (s *answerabilityGateState) recordAnswerability(status string, reason string, err error) { + s.recordAnswerabilityWithLatency(status, reason, err, time.Time{}) +} + +func (s *answerabilityGateState) recordAnswerabilityWithLatency(status string, reason string, err error, started time.Time) { + if s == nil || s.Input.Collector == nil { + return + } + errorMessage := strings.TrimSpace(s.ErrorMessage) + if err != nil { + errorMessage = err.Error() + } + data := callbacks.AnswerabilityTraceData{ + Status: status, + Reason: strings.TrimSpace(reason), + SupportingChunkIDs: append([]string(nil), s.Grade.SupportingChunkIDs...), + MissingInfo: append([]string(nil), s.Grade.MissingInfo...), + ErrorMessage: errorMessage, + } + if !started.IsZero() { + data.LatencyMs = time.Since(started).Milliseconds() + } + s.Input.Collector.SetAnswerability(data) +} diff --git a/internal/ai/runtime/executor/answerability_gate_test.go b/internal/ai/runtime/executor/answerability_gate_test.go new file mode 100644 index 0000000..36fa974 --- /dev/null +++ b/internal/ai/runtime/executor/answerability_gate_test.go @@ -0,0 +1,233 @@ +package executor + +import ( + "context" + "errors" + "strings" + "testing" + + "cs-agent/internal/ai/rag" + "cs-agent/internal/ai/runtime/internal/impl/callbacks" + "cs-agent/internal/ai/runtime/internal/impl/retrievers" + "cs-agent/internal/models" + + "github.com/cloudwego/eino/components/model" + "github.com/cloudwego/eino/schema" +) + +func TestParseAnswerabilityDecisionRejectsMalformedJSON(t *testing.T) { + _, err := parseAnswerabilityDecision(`{"answerable": true`) + + if err == nil { + t.Fatal("expected malformed JSON to be rejected") + } +} + +func TestParseAnswerabilityDecisionRejectsAnswerableWithoutSupportingChunkIDs(t *testing.T) { + _, err := parseAnswerabilityDecision(`{"answerable": true, "reason": "directly supported"}`) + + if err == nil { + t.Fatal("expected answerable decision without supporting chunks to be rejected") + } +} + +func TestParseAnswerabilityDecisionAcceptsAnswerableWithSupportingChunkIDs(t *testing.T) { + got, err := parseAnswerabilityDecision("```json\n{\"answerable\": true, \"reason\": \"directly supported\", \"supportingChunkIds\": [\" chunk-1 \", \"chunk-2\"]}\n```") + if err != nil { + t.Fatalf("parse decision failed: %v", err) + } + + if !got.Answerable { + t.Fatal("expected answerable decision") + } + if got.SupportingChunkIDs[0] != "chunk-1" || got.SupportingChunkIDs[1] != "chunk-2" { + t.Fatalf("unexpected supporting chunks: %#v", got.SupportingChunkIDs) + } +} + +func TestKnowledgeAnswerabilityGateEvaluateFallsBackOnRetrieverError(t *testing.T) { + collector := callbacks.NewRuntimeTraceCollector() + gate := newTestKnowledgeAnswerabilityGate(&fakeKnowledgeContextRetriever{ + knowledgeBaseIDs: []int64{1}, + err: errors.New("vector store unavailable"), + }, nil) + + state, err := gate.Evaluate(context.Background(), answerabilityGateInput{ + Request: newAnswerabilityGateRunInput("是否支持退款?", "1"), + Summary: &RunResult{}, + Collector: collector, + }) + if err != nil { + t.Fatalf("Evaluate returned error: %v", err) + } + + if !strings.Contains(state.FallbackReply, "建议你联系人工客服进一步确认。") { + t.Fatalf("expected human-support fallback, got %q", state.FallbackReply) + } + if collector.Data.Answerability.Status != answerabilityStatusUnanswerable { + t.Fatalf("unexpected answerability status: %q", collector.Data.Answerability.Status) + } + if collector.Data.Answerability.Reason != "knowledge retrieval failed" { + t.Fatalf("unexpected reason: %q", collector.Data.Answerability.Reason) + } +} + +func TestKnowledgeAnswerabilityGateEvaluateSkipsWhenNoKnowledgeConfigured(t *testing.T) { + collector := callbacks.NewRuntimeTraceCollector() + gate := newTestKnowledgeAnswerabilityGate(&fakeKnowledgeContextRetriever{}, nil) + + state, err := gate.Evaluate(context.Background(), answerabilityGateInput{ + Request: newAnswerabilityGateRunInput("是否支持退款?", ""), + Collector: collector, + }) + if err != nil { + t.Fatalf("Evaluate returned error: %v", err) + } + + if !state.SkipGate { + t.Fatal("expected gate to skip without knowledge") + } + if state.FallbackReply != "" { + t.Fatalf("expected no fallback when gate skips, got %q", state.FallbackReply) + } + if collector.Data.Answerability.Status != answerabilityStatusSkipped { + t.Fatalf("unexpected answerability status: %q", collector.Data.Answerability.Status) + } +} + +func TestKnowledgeAnswerabilityGateEvaluateFallsBackOnGrayZoneUnanswerableDecision(t *testing.T) { + collector := callbacks.NewRuntimeTraceCollector() + gate := newTestKnowledgeAnswerabilityGate(newAnswerabilityRetrieverWithHit(), &fakeAnswerabilityChatModel{ + response: `{"answerable": false, "reason": "retrieved snippets mention refunds but not the requested condition", "missingInfo": ["refund condition"]}`, + }) + + state, err := gate.Evaluate(context.Background(), answerabilityGateInput{ + Request: newAnswerabilityGateRunInput("满足什么条件可以退款?", "1"), + Collector: collector, + }) + if err != nil { + t.Fatalf("Evaluate returned error: %v", err) + } + + if !strings.Contains(state.FallbackReply, "建议你联系人工客服进一步确认。") { + t.Fatalf("expected human-support fallback, got %q", state.FallbackReply) + } + if collector.Data.Answerability.Status != answerabilityStatusUnanswerable { + t.Fatalf("unexpected answerability status: %q", collector.Data.Answerability.Status) + } + if collector.Data.Answerability.MissingInfo[0] != "refund condition" { + t.Fatalf("unexpected missing info: %#v", collector.Data.Answerability.MissingInfo) + } +} + +func TestKnowledgeAnswerabilityGateEvaluateAllowsAnswerableDecisionAndProducesKnowledgeInstruction(t *testing.T) { + collector := callbacks.NewRuntimeTraceCollector() + summary := &RunResult{} + chatModel := &fakeAnswerabilityChatModel{ + response: `{"answerable": true, "reason": "refund condition is directly supported", "supportingChunkIds": ["101"]}`, + } + gate := newTestKnowledgeAnswerabilityGate(newAnswerabilityRetrieverWithHit(), chatModel) + + state, err := gate.Evaluate(context.Background(), answerabilityGateInput{ + Request: newAnswerabilityGateRunInput("满足什么条件可以退款?", "1"), + Summary: summary, + Collector: collector, + }) + if err != nil { + t.Fatalf("Evaluate returned error: %v", err) + } + + if state.FallbackReply != "" { + t.Fatalf("expected answerable gate to allow, got fallback %q", state.FallbackReply) + } + if len(state.Decision.Instructions) != 1 { + t.Fatalf("expected one knowledge instruction, got %d", len(state.Decision.Instructions)) + } + if !strings.Contains(state.Decision.Instructions[0].Content, "知识库回答约束") { + t.Fatalf("unexpected instruction: %q", state.Decision.Instructions[0].Content) + } + if summary.RetrieverCount != 1 { + t.Fatalf("expected retriever count 1, got %d", summary.RetrieverCount) + } + if collector.Data.Answerability.Status != answerabilityStatusAnswerable { + t.Fatalf("unexpected answerability status: %q", collector.Data.Answerability.Status) + } + if collector.Data.Answerability.SupportingChunkIDs[0] != "101" { + t.Fatalf("unexpected supporting chunks: %#v", collector.Data.Answerability.SupportingChunkIDs) + } + if len(chatModel.input) == 0 || !strings.Contains(chatModel.input[len(chatModel.input)-1].Content, "chunkId: 101") { + t.Fatalf("expected grader prompt to include chunk id, got %#v", chatModel.input) + } +} + +func newTestKnowledgeAnswerabilityGate(retriever knowledgeContextRetriever, chatModel model.BaseChatModel) *KnowledgeAnswerabilityGate { + return &KnowledgeAnswerabilityGate{ + newRetriever: func(aiAgent models.AIAgent) knowledgeContextRetriever { + return retriever + }, + newChatModel: func(ctx context.Context, aiConfig models.AIConfig) (model.BaseChatModel, error) { + return chatModel, nil + }, + } +} + +func newAnswerabilityGateRunInput(question string, knowledgeIDs string) RunInput { + return RunInput{ + UserMessage: models.Message{Content: question}, + AIAgent: models.AIAgent{ + KnowledgeIDs: knowledgeIDs, + }, + } +} + +func newAnswerabilityRetrieverWithHit() *fakeKnowledgeContextRetriever { + return &fakeKnowledgeContextRetriever{ + knowledgeBaseIDs: []int64{1}, + result: &retrievers.KnowledgeRetrieveResult{ + KnowledgeBaseIDs: []int64{1}, + Hits: []rag.RetrieveResult{ + {KnowledgeBaseID: 1, DocumentID: 10, ChunkID: 101, Score: 0.93, Content: "购买后七天内且未使用可以退款。"}, + }, + ContextResults: []rag.RetrieveResult{ + {KnowledgeBaseID: 1, DocumentID: 10, ChunkID: 101, Score: 0.93, Content: "购买后七天内且未使用可以退款。"}, + }, + ContextText: "购买后七天内且未使用可以退款。", + }, + } +} + +type fakeKnowledgeContextRetriever struct { + knowledgeBaseIDs []int64 + result *retrievers.KnowledgeRetrieveResult + err error + lastOptions retrievers.KnowledgeRetrieveOptions + lastQuery string +} + +func (f *fakeKnowledgeContextRetriever) KnowledgeBaseIDs() []int64 { + return append([]int64(nil), f.knowledgeBaseIDs...) +} + +func (f *fakeKnowledgeContextRetriever) RetrieveContextByOptions(ctx context.Context, opts retrievers.KnowledgeRetrieveOptions, query string) (*retrievers.KnowledgeRetrieveResult, error) { + f.lastOptions = opts + f.lastQuery = query + return f.result, f.err +} + +type fakeAnswerabilityChatModel struct { + response string + err error + input []*schema.Message +} + +func (f *fakeAnswerabilityChatModel) Generate(ctx context.Context, input []*schema.Message, opts ...model.Option) (*schema.Message, error) { + f.input = input + if f.err != nil { + return nil, f.err + } + return schema.AssistantMessage(f.response, nil), nil +} + +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") +}