fix: guard runtime answers with knowledge retrieval

This commit is contained in:
mlogclub
2026-04-15 14:57:14 +08:00
parent 1ad788b3ed
commit ef73cf7e61
5 changed files with 195 additions and 4 deletions
@@ -29,21 +29,27 @@ func buildRunMessages(ctx context.Context, req RunInput, summary *RunResult, col
}
messages := make([]*schema.Message, 0, len(history.Messages)+3)
messages = append(messages, history.Messages...)
appendRetrievedContext(ctx, req, summary, collector, &messages)
decision := appendRetrievedContext(ctx, req, summary, collector, &messages)
if strings.TrimSpace(decision.FallbackReply) != "" {
if summary != nil {
summary.ReplyText = decision.FallbackReply
}
return messages
}
messages = append(messages, schema.UserMessage(strings.TrimSpace(req.UserMessage.Content)))
return messages
}
func appendRetrievedContext(ctx context.Context, req RunInput, summary *RunResult, collector *callbacks.RuntimeTraceCollector, messages *[]*schema.Message) {
func appendRetrievedContext(ctx context.Context, req RunInput, summary *RunResult, collector *callbacks.RuntimeTraceCollector, messages *[]*schema.Message) knowledgeGuardDecision {
if req.AIAgent == nil || req.UserMessage == nil || messages == nil {
return
return knowledgeGuardDecision{}
}
retriever := retrievers.NewKnowledgeRetriever(req.AIAgent)
retrieveOptions := retrievers.DefaultKnowledgeRetrieveOptions()
retrieveOptions.QueryPreview = preview(req.UserMessage.Content, 120)
retrieveResult, retrieveErr := retriever.RetrieveContextByOptions(ctx, retrieveOptions, strings.TrimSpace(req.UserMessage.Content))
if retrieveErr != nil || retrieveResult == nil {
return
return knowledgeGuardDecision{}
}
if summary != nil {
summary.RetrieverCount = len(retrieveResult.Hits)
@@ -52,7 +58,12 @@ func appendRetrievedContext(ctx context.Context, req RunInput, summary *RunResul
collector.SetRetrieverSummary(retrieveResult.TraceSummary)
collector.Data.Retriever.Items = append(collector.Data.Retriever.Items, retrieveResult.TraceItems...)
}
decision := buildKnowledgeGuardDecision(req.AIAgent, retrieveResult)
if len(decision.Instructions) > 0 {
*messages = append(*messages, decision.Instructions...)
}
if strings.TrimSpace(retrieveResult.ContextText) != "" {
*messages = append(*messages, schema.SystemMessage(retrieveResult.ContextText))
}
return decision
}
@@ -0,0 +1,60 @@
package executor
import (
"strings"
"cs-agent/internal/ai/runtime/internal/impl/retrievers"
"cs-agent/internal/models"
"cs-agent/internal/pkg/enums"
"github.com/cloudwego/eino/schema"
)
type knowledgeGuardDecision struct {
FallbackReply string
Instructions []*schema.Message
}
func buildKnowledgeGuardDecision(aiAgent *models.AIAgent, retrieveResult *retrievers.KnowledgeRetrieveResult) knowledgeGuardDecision {
if aiAgent == nil || retrieveResult == nil || len(retrieveResult.KnowledgeBaseIDs) == 0 {
return knowledgeGuardDecision{}
}
fallbackReply := resolveKnowledgeFallbackReply(aiAgent, retrieveResult.FallbackMode)
if len(retrieveResult.Hits) == 0 {
return knowledgeGuardDecision{FallbackReply: fallbackReply}
}
instruction := buildKnowledgeRuntimeInstruction(retrieveResult.AnswerMode, fallbackReply)
if instruction == "" {
return knowledgeGuardDecision{}
}
return knowledgeGuardDecision{
Instructions: []*schema.Message{schema.SystemMessage(instruction)},
}
}
func resolveKnowledgeFallbackReply(aiAgent *models.AIAgent, fallbackMode enums.KnowledgeFallbackMode) string {
if aiAgent != nil {
if reply := strings.TrimSpace(aiAgent.FallbackMessage); reply != "" {
return reply
}
}
switch fallbackMode {
case enums.KnowledgeFallbackModeSuggestRetry:
return "当前知识库里没有找到足够明确的信息,你可以换个更具体的问法再试一次。"
case enums.KnowledgeFallbackModeTransferHuman:
return "当前知识库里没有找到足够明确的信息,建议转人工进一步处理。"
default:
return "当前知识库暂无明确信息。"
}
}
func buildKnowledgeRuntimeInstruction(answerMode enums.KnowledgeAnswerMode, fallbackReply string) string {
fallbackReply = strings.TrimSpace(fallbackReply)
if fallbackReply == "" {
fallbackReply = "当前知识库暂无明确信息。"
}
if answerMode == enums.KnowledgeAnswerModeAssist {
return "知识库回答约束:优先依据后续提供的知识片段回答,可以做轻度归纳,但不要编造片段中未提供的事实。若知识片段不足以直接支持答案,必须明确回复:" + fallbackReply
}
return "知识库回答约束:本轮只能依据后续提供的知识片段回答,不得使用模型常识补充未提供的事实。若知识片段不足以支持回答,必须明确回复:" + fallbackReply
}
@@ -0,0 +1,69 @@
package executor
import (
"strings"
"testing"
"cs-agent/internal/ai/rag"
"cs-agent/internal/ai/runtime/internal/impl/retrievers"
"cs-agent/internal/models"
"cs-agent/internal/pkg/enums"
)
func TestBuildKnowledgeGuardDecisionFallsBackWhenKnowledgeMisses(t *testing.T) {
agent := newKnowledgeGuardAgentFixture()
decision := buildKnowledgeGuardDecision(&agent, &retrievers.KnowledgeRetrieveResult{
KnowledgeBaseIDs: []int64{1},
FallbackMode: enums.KnowledgeFallbackModeSuggestRetry,
})
if decision.FallbackReply != "当前知识库里没有找到足够明确的信息,你可以换个更具体的问法再试一次。" {
t.Fatalf("unexpected fallback reply: %q", decision.FallbackReply)
}
if len(decision.Instructions) != 0 {
t.Fatalf("expected no instructions on miss, got %d", len(decision.Instructions))
}
}
func TestBuildKnowledgeGuardDecisionUsesAgentFallbackMessage(t *testing.T) {
agent := newKnowledgeGuardAgentFixture()
agent.FallbackMessage = "请联系人工客服"
decision := buildKnowledgeGuardDecision(&agent, &retrievers.KnowledgeRetrieveResult{
KnowledgeBaseIDs: []int64{1},
FallbackMode: enums.KnowledgeFallbackModeNoAnswer,
})
if decision.FallbackReply != "请联系人工客服" {
t.Fatalf("expected agent fallback message, got %q", decision.FallbackReply)
}
}
func TestBuildKnowledgeGuardDecisionInjectsStrictInstructionOnHit(t *testing.T) {
agent := newKnowledgeGuardAgentFixture()
decision := buildKnowledgeGuardDecision(&agent, &retrievers.KnowledgeRetrieveResult{
KnowledgeBaseIDs: []int64{1},
Hits: []rag.RetrieveResult{
{KnowledgeBaseID: 1, Score: 0.88},
},
AnswerMode: enums.KnowledgeAnswerModeStrict,
FallbackMode: enums.KnowledgeFallbackModeNoAnswer,
})
if decision.FallbackReply != "" {
t.Fatalf("expected no fallback reply on hit, got %q", decision.FallbackReply)
}
if len(decision.Instructions) != 1 {
t.Fatalf("expected one instruction, got %d", len(decision.Instructions))
}
content := decision.Instructions[0].Content
if !strings.Contains(content, "只能依据后续提供的知识片段回答") {
t.Fatalf("unexpected strict instruction: %q", content)
}
if !strings.Contains(content, "当前知识库暂无明确信息。") {
t.Fatalf("expected fallback text in instruction, got %q", content)
}
}
func newKnowledgeGuardAgentFixture() models.AIAgent {
return models.AIAgent{}
}
+9
View File
@@ -118,6 +118,15 @@ func (s *Service) ExecuteRun(ctx context.Context, req RunInput) (*RunResult, err
return summary, fmt.Errorf("%s", summary.ErrorMessage)
}
messages := buildRunMessages(ctx, req, summary, collector)
if strings.TrimSpace(summary.ReplyText) != "" {
summary.Status = "completed"
summary.ModelName = req.AIConfig.ModelName
collector.Data.Status = summary.Status
collector.Data.Output.ReplyText = summary.ReplyText
collector.Data.Output.FinishReason = summary.Status
summary.TraceData = collector.Marshal()
return summary, nil
}
collector.Data.Interrupt.CheckPointID = checkPointID
consumeAgentEvents(runner.Run(ctx, messages, buildRunOptions(checkPointID)...), summary, collector, tooling.toolDefsByModelName)
summary.ModelName = req.AIConfig.ModelName