Merge branch 'codex/knowledge-runtime-guard'
This commit is contained in:
@@ -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{}
|
||||
}
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user