fix: harden rag answerability gate
This commit is contained in:
@@ -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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user