Merge branch 'codex/knowledge-runtime-guard'

This commit is contained in:
mlogclub
2026-04-15 14:57:30 +08:00
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
@@ -44,6 +44,9 @@ type KnowledgeRetrieveResult struct {
Hits []rag.RetrieveResult
ContextResults []rag.RetrieveResult
ContextText string
TopScore float64
AnswerMode enums.KnowledgeAnswerMode
FallbackMode enums.KnowledgeFallbackMode
Trace *rag.RetrieveTrace
TraceItems []callbacks.RetrieverTraceItem
TraceSummary callbacks.RetrieverTraceSummary
@@ -126,6 +129,8 @@ func (r *KnowledgeRetriever) RetrieveContextByOptions(ctx context.Context, opts
ret.ContextResults = rag.Retrieve.SelectContextResults(results, contextMaxTokens)
ret.ContextResults = limitContextResults(ret.ContextResults, maxContextItems)
ret.ContextText = strings.TrimSpace(buildContextText(ret.ContextResults))
ret.TopScore = resolveTopScore(results)
ret.AnswerMode, ret.FallbackMode = resolveRuntimeAnswerSettings(knowledgeBaseIDs, results)
ret.TraceItems = buildRetrieverTraceItems(queryPreview, results, trace)
ret.TraceSummary = buildRetrieverTraceSummary(ret.Options, ret.Policies, ret.ContextResults, results, trace)
return ret, nil
@@ -148,6 +153,13 @@ func buildContextText(results []rag.RetrieveResult) string {
return strings.TrimSpace(rag.Retrieve.BuildContext(context.Background(), results, 1<<30))
}
func resolveTopScore(results []rag.RetrieveResult) float64 {
if len(results) == 0 {
return 0
}
return float64(results[0].Score)
}
func (r *KnowledgeRetriever) resolvePolicies(knowledgeBaseIDs []int64, opts KnowledgeRetrieveOptions) []KnowledgeBaseRetrievePolicy {
if len(knowledgeBaseIDs) == 0 {
return nil
@@ -179,6 +191,36 @@ func (r *KnowledgeRetriever) resolvePolicies(knowledgeBaseIDs []int64, opts Know
return ret
}
func resolveRuntimeAnswerSettings(knowledgeBaseIDs []int64, results []rag.RetrieveResult) (enums.KnowledgeAnswerMode, enums.KnowledgeFallbackMode) {
knowledgeBases := loadRuntimeKnowledgeBases(knowledgeBaseIDs)
if len(knowledgeBases) == 0 {
return enums.KnowledgeAnswerModeStrict, enums.KnowledgeFallbackModeNoAnswer
}
if len(results) > 0 {
if knowledgeBase, ok := knowledgeBases[results[0].KnowledgeBaseID]; ok {
return normalizeRuntimeAnswerSettings(knowledgeBase)
}
}
for _, knowledgeBaseID := range knowledgeBaseIDs {
if knowledgeBase, ok := knowledgeBases[knowledgeBaseID]; ok {
return normalizeRuntimeAnswerSettings(knowledgeBase)
}
}
return enums.KnowledgeAnswerModeStrict, enums.KnowledgeFallbackModeNoAnswer
}
func normalizeRuntimeAnswerSettings(knowledgeBase models.KnowledgeBase) (enums.KnowledgeAnswerMode, enums.KnowledgeFallbackMode) {
answerMode := enums.KnowledgeAnswerMode(knowledgeBase.AnswerMode)
if answerMode == 0 {
answerMode = enums.KnowledgeAnswerModeStrict
}
fallbackMode := enums.KnowledgeFallbackMode(knowledgeBase.FallbackMode)
if fallbackMode == 0 {
fallbackMode = enums.KnowledgeFallbackModeNoAnswer
}
return answerMode, fallbackMode
}
func loadRuntimeKnowledgeBases(ids []int64) map[int64]models.KnowledgeBase {
if len(ids) == 0 {
return nil