feat: wire answerability gate into ai runtime

This commit is contained in:
mlogclub
2026-05-02 14:21:49 +08:00
parent 0230e05dce
commit 8f968fe3ca
3 changed files with 79 additions and 27 deletions
@@ -186,6 +186,46 @@ func TestKnowledgeAnswerabilityGateEvaluateAllowsAnswerableDecisionAndProducesKn
}
}
func TestBuildRunMessagesReturnsFallbackWhenGateRejectsGrayZoneQuestion(t *testing.T) {
summary := &RunResult{}
question := "满足什么条件可以退款?"
gate := newTestKnowledgeAnswerabilityGate(newAnswerabilityRetrieverWithHit(), &fakeAnswerabilityChatModel{
response: `{"answerable": false, "reason": "retrieved snippets mention refunds but not the requested condition", "missingInfo": ["refund condition"]}`,
})
messages := buildRunMessages(context.Background(), newAnswerabilityGateRunInput(question, "1"), summary, nil, gate)
if !strings.Contains(summary.ReplyText, "建议你联系人工客服进一步确认。") {
t.Fatalf("expected human-support fallback, got %q", summary.ReplyText)
}
if messagesContainContent(messages, question) {
t.Fatalf("expected returned messages to omit current user message, got %#v", messages)
}
}
func TestBuildRunMessagesInjectsKnowledgeWhenGateAllows(t *testing.T) {
summary := &RunResult{}
question := "满足什么条件可以退款?"
gate := newTestKnowledgeAnswerabilityGate(newAnswerabilityRetrieverWithHit(), &fakeAnswerabilityChatModel{
response: `{"answerable": true, "reason": "refund condition is directly supported", "supportingChunkIds": ["101"]}`,
})
messages := buildRunMessages(context.Background(), newAnswerabilityGateRunInput(question, "1"), summary, nil, gate)
if summary.ReplyText != "" {
t.Fatalf("expected no fallback, got %q", summary.ReplyText)
}
if !messagesContainContent(messages, "知识库回答约束") {
t.Fatalf("expected knowledge instruction in messages, got %#v", messages)
}
if !messagesContainContent(messages, "购买后七天内且未使用可以退款。") {
t.Fatalf("expected knowledge context in messages, got %#v", messages)
}
if !messagesContainContent(messages, question) {
t.Fatalf("expected current user message in messages, got %#v", messages)
}
}
func newTestKnowledgeAnswerabilityGate(retriever knowledgeContextRetriever, chatModel model.BaseChatModel) *KnowledgeAnswerabilityGate {
return &KnowledgeAnswerabilityGate{
newRetriever: func(aiAgent models.AIAgent) knowledgeContextRetriever {
@@ -257,3 +297,12 @@ func (f *fakeAnswerabilityChatModel) Generate(ctx context.Context, input []*sche
func (f *fakeAnswerabilityChatModel) Stream(ctx context.Context, input []*schema.Message, opts ...model.Option) (*schema.StreamReader[*schema.Message], error) {
return nil, errors.New("stream is not implemented in fakeAnswerabilityChatModel")
}
func messagesContainContent(messages []*schema.Message, text string) bool {
for _, message := range messages {
if message != nil && strings.Contains(message.Content, text) {
return true
}
}
return false
}
@@ -6,13 +6,12 @@ import (
"cs-agent/internal/ai/runtime/internal/impl/adapter"
"cs-agent/internal/ai/runtime/internal/impl/callbacks"
"cs-agent/internal/ai/runtime/internal/impl/retrievers"
"cs-agent/internal/pkg/utils"
"github.com/cloudwego/eino/schema"
)
func buildRunMessages(ctx context.Context, req RunInput, summary *RunResult, collector *callbacks.RuntimeTraceCollector) []*schema.Message {
func buildRunMessages(ctx context.Context, req RunInput, summary *RunResult, collector *callbacks.RuntimeTraceCollector, gate *KnowledgeAnswerabilityGate) []*schema.Message {
history := adapter.BuildHistoryMessages(req.Conversation.ID, req.UserMessage.ID, 12)
if summary != nil {
summary.HistoryMessageCount = len(history.Messages)
@@ -24,7 +23,7 @@ func buildRunMessages(ctx context.Context, req RunInput, summary *RunResult, col
}
messages := make([]*schema.Message, 0, len(history.Messages)+3)
messages = append(messages, history.Messages...)
decision := appendRetrievedContext(ctx, req, summary, collector, &messages)
decision := appendRetrievedContext(ctx, req, summary, collector, gate, &messages)
if strings.TrimSpace(decision.FallbackReply) != "" {
if summary != nil {
summary.ReplyText = decision.FallbackReply
@@ -35,30 +34,32 @@ func buildRunMessages(ctx context.Context, req RunInput, summary *RunResult, col
return messages
}
func appendRetrievedContext(ctx context.Context, req RunInput, summary *RunResult, collector *callbacks.RuntimeTraceCollector, messages *[]*schema.Message) knowledgeGuardDecision {
func appendRetrievedContext(ctx context.Context, req RunInput, summary *RunResult, collector *callbacks.RuntimeTraceCollector, gate *KnowledgeAnswerabilityGate, messages *[]*schema.Message) knowledgeGuardDecision {
if messages == nil {
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 buildKnowledgeUnavailableDecision(req.AIAgent, retriever.KnowledgeBaseIDs())
if gate == nil {
gate = NewKnowledgeAnswerabilityGate()
}
if summary != nil {
summary.RetrieverCount = len(retrieveResult.Hits)
state, err := gate.Evaluate(ctx, answerabilityGateInput{
Request: req,
Summary: summary,
Collector: collector,
Messages: append([]*schema.Message(nil), (*messages)...),
})
if err != nil || state == nil {
decision := buildKnowledgeUnavailableDecision(req.AIAgent, utils.SplitInt64s(req.AIAgent.KnowledgeIDs))
if strings.TrimSpace(decision.FallbackReply) != "" {
decision.FallbackReply = resolveKnowledgeHumanSupportFallback(req.AIAgent)
}
return decision
}
if collector != nil {
collector.SetRetrieverSummary(retrieveResult.TraceSummary)
collector.AddRetrieverItems(retrieveResult.TraceItems)
if strings.TrimSpace(state.FallbackReply) != "" {
return knowledgeGuardDecision{FallbackReply: state.FallbackReply}
}
decision := buildKnowledgeGuardDecision(req.AIAgent, retrieveResult)
if len(decision.Instructions) > 0 {
*messages = append(*messages, decision.Instructions...)
if state.SkipGate {
return knowledgeGuardDecision{}
}
if strings.TrimSpace(retrieveResult.ContextText) != "" {
*messages = append(*messages, schema.SystemMessage(retrieveResult.ContextText))
}
return decision
*messages = append((*messages)[:0], state.Input.Messages...)
return state.Decision
}
+7 -5
View File
@@ -13,14 +13,16 @@ import (
)
type Service struct {
agentFactory *factory.AgentFactory
runnerFactory *factory.RunnerFactory
agentFactory *factory.AgentFactory
runnerFactory *factory.RunnerFactory
answerabilityGate *KnowledgeAnswerabilityGate
}
func NewService() *Service {
return &Service{
agentFactory: factory.NewAgentFactory(),
runnerFactory: factory.NewRunnerFactory(),
agentFactory: factory.NewAgentFactory(),
runnerFactory: factory.NewRunnerFactory(),
answerabilityGate: NewKnowledgeAnswerabilityGate(),
}
}
@@ -84,7 +86,7 @@ func (s *Service) ExecuteRun(ctx context.Context, req RunInput) (*RunResult, err
summary.TraceData = collector.Marshal()
return summary, fmt.Errorf("%s", summary.ErrorMessage)
}
messages := buildRunMessages(ctx, req, summary, collector)
messages := buildRunMessages(ctx, req, summary, collector, s.answerabilityGate)
if strings.TrimSpace(summary.ReplyText) != "" {
summary.Status = "completed"
summary.ModelName = req.AIConfig.ModelName