feat: implement rag answerability gate graph

This commit is contained in:
mlogclub
2026-05-02 14:06:23 +08:00
parent 269042b9d4
commit 9e30167eaf
2 changed files with 576 additions and 0 deletions
@@ -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 arrayanswerable 为 true 时必须至少包含一个直接支持的 chunk id
- missingInfo: string arrayanswerable 为 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)
}
@@ -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")
}