fix: harden rag answerability gate

This commit is contained in:
mlogclub
2026-05-02 14:14:31 +08:00
parent 9e30167eaf
commit 0230e05dce
5 changed files with 68 additions and 5 deletions
@@ -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 {
@@ -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{
@@ -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 {
@@ -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()
@@ -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)
}
}