From 0230e05dce2f1d2b77ee2fcb30f6445c868ba2f1 Mon Sep 17 00:00:00 2001 From: mlogclub Date: Sat, 2 May 2026 14:14:31 +0800 Subject: [PATCH] fix: harden rag answerability gate --- .../ai/runtime/executor/answerability_gate.go | 16 +++++++++--- .../executor/answerability_gate_test.go | 26 +++++++++++++++++++ .../ai/runtime/executor/context_builders.go | 2 +- .../impl/callbacks/runlog_callback.go | 9 +++++++ .../impl/callbacks/runlog_callback_test.go | 20 ++++++++++++++ 5 files changed, 68 insertions(+), 5 deletions(-) create mode 100644 internal/ai/runtime/internal/impl/callbacks/runlog_callback_test.go diff --git a/internal/ai/runtime/executor/answerability_gate.go b/internal/ai/runtime/executor/answerability_gate.go index 1ef82fd..1b5a865 100644 --- a/internal/ai/runtime/executor/answerability_gate.go +++ b/internal/ai/runtime/executor/answerability_gate.go @@ -11,6 +11,7 @@ import ( "cs-agent/internal/ai/runtime/internal/impl/factory" "cs-agent/internal/ai/runtime/internal/impl/retrievers" "cs-agent/internal/models" + "cs-agent/internal/pkg/utils" "github.com/cloudwego/eino/components/model" "github.com/cloudwego/eino/components/prompt" @@ -178,10 +179,17 @@ func (g *KnowledgeAnswerabilityGate) retrieveKnowledge(ctx context.Context, stat } gate := g.withDefaults() req := state.Input.Request + configuredKnowledgeIDs := utils.SplitInt64s(req.AIAgent.KnowledgeIDs) retriever := gate.newRetriever(req.AIAgent) if retriever == nil { - state.SkipGate = true - state.recordAnswerability(answerabilityStatusSkipped, "knowledge retriever unavailable", nil) + state.KnowledgeIDs = append([]int64(nil), configuredKnowledgeIDs...) + if len(configuredKnowledgeIDs) == 0 { + state.SkipGate = true + state.recordAnswerability(answerabilityStatusSkipped, "no knowledge configured", nil) + return state, nil + } + state.FallbackReply = resolveKnowledgeHumanSupportFallback(req.AIAgent) + state.recordAnswerability(answerabilityStatusUnanswerable, "knowledge retriever unavailable", nil) return state, nil } knowledgeIDs := retriever.KnowledgeBaseIDs() @@ -212,7 +220,7 @@ func (g *KnowledgeAnswerabilityGate) retrieveKnowledge(ctx context.Context, stat } 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...) + state.Input.Collector.AddRetrieverItems(result.TraceItems) } if result == nil || len(result.Hits) == 0 || strings.TrimSpace(result.ContextText) == "" { state.FallbackReply = resolveKnowledgeHumanSupportFallback(req.AIAgent) @@ -230,7 +238,7 @@ func (g *KnowledgeAnswerabilityGate) gradeAnswerability(ctx context.Context, sta return state, nil } gate := g.withDefaults() - started := time.Now() + started := gate.now() req := state.Input.Request modelInstance, err := gate.newChatModel(ctx, req.AIConfig) if err != nil { diff --git a/internal/ai/runtime/executor/answerability_gate_test.go b/internal/ai/runtime/executor/answerability_gate_test.go index 36fa974..4520a9a 100644 --- a/internal/ai/runtime/executor/answerability_gate_test.go +++ b/internal/ai/runtime/executor/answerability_gate_test.go @@ -95,6 +95,32 @@ func TestKnowledgeAnswerabilityGateEvaluateSkipsWhenNoKnowledgeConfigured(t *tes } } +func TestKnowledgeAnswerabilityGateEvaluateFallsBackWhenConfiguredRetrieverUnavailable(t *testing.T) { + collector := callbacks.NewRuntimeTraceCollector() + gate := newTestKnowledgeAnswerabilityGate(nil, nil) + + state, err := gate.Evaluate(context.Background(), answerabilityGateInput{ + Request: newAnswerabilityGateRunInput("是否支持退款?", "1"), + Collector: collector, + }) + if err != nil { + t.Fatalf("Evaluate returned error: %v", err) + } + + if state.SkipGate { + t.Fatal("expected configured knowledge to fail closed, not skip") + } + 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 retriever unavailable" { + t.Fatalf("unexpected reason: %q", collector.Data.Answerability.Reason) + } +} + func TestKnowledgeAnswerabilityGateEvaluateFallsBackOnGrayZoneUnanswerableDecision(t *testing.T) { collector := callbacks.NewRuntimeTraceCollector() gate := newTestKnowledgeAnswerabilityGate(newAnswerabilityRetrieverWithHit(), &fakeAnswerabilityChatModel{ diff --git a/internal/ai/runtime/executor/context_builders.go b/internal/ai/runtime/executor/context_builders.go index 8ef1fcf..bf0ca08 100644 --- a/internal/ai/runtime/executor/context_builders.go +++ b/internal/ai/runtime/executor/context_builders.go @@ -51,7 +51,7 @@ func appendRetrievedContext(ctx context.Context, req RunInput, summary *RunResul } if collector != nil { collector.SetRetrieverSummary(retrieveResult.TraceSummary) - collector.Data.Retriever.Items = append(collector.Data.Retriever.Items, retrieveResult.TraceItems...) + collector.AddRetrieverItems(retrieveResult.TraceItems) } decision := buildKnowledgeGuardDecision(req.AIAgent, retrieveResult) if len(decision.Instructions) > 0 { diff --git a/internal/ai/runtime/internal/impl/callbacks/runlog_callback.go b/internal/ai/runtime/internal/impl/callbacks/runlog_callback.go index 7a45a88..ff5329a 100644 --- a/internal/ai/runtime/internal/impl/callbacks/runlog_callback.go +++ b/internal/ai/runtime/internal/impl/callbacks/runlog_callback.go @@ -106,6 +106,15 @@ func (c *RuntimeTraceCollector) SetRetrieverSummary(summary RetrieverTraceSummar c.Data.Retriever.Policies = append([]RetrieverPolicyTraceItem(nil), summary.Policies...) } +func (c *RuntimeTraceCollector) AddRetrieverItems(items []RetrieverTraceItem) { + if len(items) == 0 { + return + } + c.mu.Lock() + defer c.mu.Unlock() + c.Data.Retriever.Items = append(c.Data.Retriever.Items, items...) +} + func (c *RuntimeTraceCollector) SetAnswerability(data AnswerabilityTraceData) { c.mu.Lock() defer c.mu.Unlock() diff --git a/internal/ai/runtime/internal/impl/callbacks/runlog_callback_test.go b/internal/ai/runtime/internal/impl/callbacks/runlog_callback_test.go new file mode 100644 index 0000000..7e7ef35 --- /dev/null +++ b/internal/ai/runtime/internal/impl/callbacks/runlog_callback_test.go @@ -0,0 +1,20 @@ +package callbacks + +import "testing" + +func TestRuntimeTraceCollectorAddRetrieverItems(t *testing.T) { + collector := NewRuntimeTraceCollector() + + collector.AddRetrieverItems([]RetrieverTraceItem{ + {KnowledgeBaseID: 1, DocumentID: 10, DocumentTitle: "doc-1"}, + {KnowledgeBaseID: 2, DocumentID: 20, DocumentTitle: "doc-2"}, + }) + collector.AddRetrieverItems(nil) + + if got := len(collector.Data.Retriever.Items); got != 2 { + t.Fatalf("expected two retriever items, got %d", got) + } + if collector.Data.Retriever.Items[0].DocumentTitle != "doc-1" || collector.Data.Retriever.Items[1].DocumentTitle != "doc-2" { + t.Fatalf("unexpected retriever items: %#v", collector.Data.Retriever.Items) + } +}