feat: wire answerability gate into ai runtime
This commit is contained in:
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user