diff --git a/internal/ai/runtime/executor/context_builders.go b/internal/ai/runtime/executor/context_builders.go index f6646bd..a194a65 100644 --- a/internal/ai/runtime/executor/context_builders.go +++ b/internal/ai/runtime/executor/context_builders.go @@ -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 } diff --git a/internal/ai/runtime/executor/knowledge_guard.go b/internal/ai/runtime/executor/knowledge_guard.go new file mode 100644 index 0000000..9ecbe74 --- /dev/null +++ b/internal/ai/runtime/executor/knowledge_guard.go @@ -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 +} diff --git a/internal/ai/runtime/executor/knowledge_guard_test.go b/internal/ai/runtime/executor/knowledge_guard_test.go new file mode 100644 index 0000000..ffc80ac --- /dev/null +++ b/internal/ai/runtime/executor/knowledge_guard_test.go @@ -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{} +} diff --git a/internal/ai/runtime/executor/service.go b/internal/ai/runtime/executor/service.go index 265d064..cf4eb5d 100644 --- a/internal/ai/runtime/executor/service.go +++ b/internal/ai/runtime/executor/service.go @@ -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 diff --git a/internal/ai/runtime/internal/impl/retrievers/knowledge_retriever.go b/internal/ai/runtime/internal/impl/retrievers/knowledge_retriever.go index eca5317..77eba1d 100644 --- a/internal/ai/runtime/internal/impl/retrievers/knowledge_retriever.go +++ b/internal/ai/runtime/internal/impl/retrievers/knowledge_retriever.go @@ -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