Refactor AI skill routing and introduce reply handling services

- Moved the skill routing logic from matcher.go to a new router.go file for better organization.
- Implemented replyCommitService to handle sending AI replies and managing reply rounds.
- Added replyInterruptService to manage conversation interrupts and resume handling.
- Created replyRunLogService to log AI reply actions and their outcomes.
- Introduced helper functions for building conversation interrupts and resolving prompts.
- Added unit tests for the new services and functions to ensure correctness.
- Removed unused code and optimized imports in matcher.go.
This commit is contained in:
mlogclub
2026-04-13 17:17:13 +08:00
parent 93cda6ae82
commit 3c0abaaedc
9 changed files with 942 additions and 650 deletions
@@ -0,0 +1,68 @@
package runtime
import (
"fmt"
"strings"
"time"
"cs-agent/internal/models"
"cs-agent/internal/pkg/dto"
"cs-agent/internal/pkg/enums"
"cs-agent/internal/repositories"
svc "cs-agent/internal/services"
"github.com/mlogclub/simple/sqls"
)
type replyCommitService struct{}
func newReplyCommitService() *replyCommitService {
return &replyCommitService{}
}
func (s *replyCommitService) SendAIReply(conversation models.Conversation, message models.Message, aiAgent models.AIAgent,
replyText string, trace *aiReplyTraceData, clientPrefix string) (*models.Message, error) {
replyText = strings.TrimSpace(replyText)
if replyText == "" {
return nil, nil
}
commitStartedAt := time.Now()
replyMessage, err := svc.MessageService.SendAIMessage(
conversation.ID,
aiAgent.ID,
fmt.Sprintf("%s_%d", strings.TrimSpace(clientPrefix), message.ID),
enums.IMMessageTypeText,
replyText,
"",
s.buildAIPrincipal(aiAgent),
)
if trace != nil {
trace.CommitMs = time.Since(commitStartedAt).Milliseconds()
trace.ReplySent = err == nil && replyMessage != nil
if replyMessage != nil {
trace.ReplyMessageID = replyMessage.ID
}
}
return replyMessage, err
}
func (s *replyCommitService) IncrementAIReplyRounds(conversationID int64, nextRounds int, aiAgentName string) error {
return repositories.ConversationRepository.Updates(sqls.DB(), conversationID, map[string]any{
"ai_reply_rounds": nextRounds,
"update_user_id": 0,
"update_user_name": strings.TrimSpace(aiAgentName),
"updated_at": time.Now(),
})
}
func (s *replyCommitService) buildAIPrincipal(aiAgent models.AIAgent) *dto.AuthPrincipal {
username := "AI"
if strings.TrimSpace(aiAgent.Name) != "" {
username = aiAgent.Name
}
return &dto.AuthPrincipal{
UserID: 0,
Username: username,
Nickname: username,
}
}
@@ -0,0 +1,83 @@
package runtime
import (
"encoding/json"
"strings"
"time"
"cs-agent/internal/models"
svc "cs-agent/internal/services"
)
func buildConversationInterrupt(conversation models.Conversation, message models.Message, aiAgent models.AIAgent, summary *Summary) *models.ConversationInterrupt {
if summary == nil {
return nil
}
now := time.Now()
item := svc.ConversationInterruptService.GetByCheckPointID(summary.CheckPointID)
if item == nil {
item = &models.ConversationInterrupt{
CheckPointID: summary.CheckPointID,
CreatedAt: now,
}
}
item.ConversationID = conversation.ID
item.AIAgentID = aiAgent.ID
item.SourceMessageID = message.ID
item.InterruptID = firstInterruptID(summary)
item.InterruptType = firstInterruptType(summary)
item.Status = "pending"
item.PromptText = resolveInterruptPrompt(summary)
item.UpdatedAt = now
return item
}
func resolveInterruptPrompt(summary *Summary) string {
if summary == nil || len(summary.Interrupts) == 0 {
return "请继续补充信息后再试。"
}
if prompt := extractInterruptMessage(summary.Interrupts[0].InfoPreview); prompt != "" {
return prompt
}
if prompt := strings.TrimSpace(summary.Interrupts[0].InfoPreview); prompt != "" {
return prompt
}
return "请继续补充信息后再试。"
}
func extractInterruptMessage(infoPreview string) string {
infoPreview = strings.TrimSpace(infoPreview)
if infoPreview == "" {
return ""
}
payload := make(map[string]any)
if err := json.Unmarshal([]byte(infoPreview), &payload); err != nil {
return ""
}
if message, ok := payload["message"].(string); ok {
return strings.TrimSpace(message)
}
return ""
}
func firstInterruptID(summary *Summary) string {
if summary == nil || len(summary.Interrupts) == 0 {
return ""
}
return strings.TrimSpace(summary.Interrupts[0].ID)
}
func firstInterruptType(summary *Summary) string {
if summary == nil || len(summary.Interrupts) == 0 {
return ""
}
return strings.TrimSpace(summary.Interrupts[0].Type)
}
func isCheckpointMissingError(err error) bool {
if err == nil {
return false
}
message := strings.ToLower(strings.TrimSpace(err.Error()))
return strings.Contains(message, "failed to load from checkpoint") && strings.Contains(message, "not exist")
}
@@ -0,0 +1,113 @@
package runtime
import (
"context"
"strings"
"cs-agent/internal/ai/runtime/graphs"
"cs-agent/internal/models"
svc "cs-agent/internal/services"
)
type replyInterruptService struct{}
func newReplyInterruptService() *replyInterruptService {
return &replyInterruptService{}
}
func (s *replyInterruptService) ResumePendingInterrupt(ctx context.Context, owner *aiReplyService, conversation models.Conversation, message models.Message, aiAgent models.AIAgent,
pendingInterrupt *models.ConversationInterrupt, trace *aiReplyTraceData, summaryRef **Summary) error {
if pendingInterrupt == nil || owner == nil || owner.executor == nil {
return nil
}
summary, err := owner.executor.ResumePendingInterrupt(ctx, conversation, message, aiAgent, pendingInterrupt, trace)
*summaryRef = summary
if err != nil {
if isCheckpointMissingError(err) {
summary = expiredInterruptSummary()
*summaryRef = summary
trace.Status = "interrupt_expired"
trace.FinalAction = "expired"
replyMessage, expireErr := owner.commit.SendAIReply(conversation, message, aiAgent, summary.ReplyText, trace, "ai_interrupt_expired")
if expireErr != nil {
return expireErr
}
if err := owner.commit.IncrementAIReplyRounds(conversation.ID, conversation.AIReplyRounds+1, aiAgent.Name); err != nil {
return err
}
lastResumeMessageID := int64(0)
if replyMessage != nil {
lastResumeMessageID = replyMessage.ID
}
if expireMarkErr := svc.ConversationInterruptService.MarkExpired(pendingInterrupt.ID, lastResumeMessageID); expireMarkErr != nil {
return expireMarkErr
}
return nil
}
return err
}
if summary != nil && summary.Interrupted {
return s.HandleInterruptedResume(owner, conversation, message, aiAgent, pendingInterrupt, summary, trace)
}
if summary != nil && strings.TrimSpace(summary.ReplyText) != "" {
replyMessage, err := owner.commit.SendAIReply(conversation, message, aiAgent, summary.ReplyText, trace, "ai_resume")
if err != nil {
return err
}
if err := owner.commit.IncrementAIReplyRounds(conversation.ID, conversation.AIReplyRounds+1, aiAgent.Name); err != nil {
return err
}
replyMessageID := int64(0)
if replyMessage != nil {
replyMessageID = replyMessage.ID
}
if graphs.IsCancellationReply(summary.ReplyText) {
return svc.ConversationInterruptService.MarkCancelled(pendingInterrupt.ID, replyMessageID)
}
return svc.ConversationInterruptService.MarkResolved(pendingInterrupt.ID, replyMessageID)
}
return svc.ConversationInterruptService.MarkResolved(pendingInterrupt.ID, 0)
}
func (s *replyInterruptService) HandleInterruptedSummary(owner *aiReplyService, conversation models.Conversation, message models.Message, aiAgent models.AIAgent,
summary *Summary, trace *aiReplyTraceData) error {
if owner == nil {
return nil
}
pending := buildConversationInterrupt(conversation, message, aiAgent, summary)
if err := svc.ConversationInterruptService.CreateOrUpdatePending(pending); err != nil {
return err
}
pending = svc.ConversationInterruptService.GetByCheckPointID(summary.CheckPointID)
replyText := resolveInterruptPrompt(summary)
replyMessage, err := owner.commit.SendAIReply(conversation, message, aiAgent, replyText, trace, "ai_interrupt")
if err != nil {
return err
}
if err := owner.commit.IncrementAIReplyRounds(conversation.ID, conversation.AIReplyRounds+1, aiAgent.Name); err != nil {
return err
}
if replyMessage != nil && pending != nil {
return svc.ConversationInterruptService.MarkPendingAgain(pending.ID, pending.InterruptID, replyText, replyMessage.ID)
}
return nil
}
func (s *replyInterruptService) HandleInterruptedResume(owner *aiReplyService, conversation models.Conversation, message models.Message, aiAgent models.AIAgent,
pendingInterrupt *models.ConversationInterrupt, summary *Summary, trace *aiReplyTraceData) error {
if pendingInterrupt == nil || owner == nil {
return nil
}
replyText := resolveInterruptPrompt(summary)
replyMessage, err := owner.commit.SendAIReply(conversation, message, aiAgent, replyText, trace, "ai_interrupt_resume")
if err != nil {
return err
}
if err := owner.commit.IncrementAIReplyRounds(conversation.ID, conversation.AIReplyRounds+1, aiAgent.Name); err != nil {
return err
}
if replyMessage != nil {
return svc.ConversationInterruptService.MarkPendingAgain(pendingInterrupt.ID, firstInterruptID(summary), replyText, replyMessage.ID)
}
return nil
}
+342
View File
@@ -0,0 +1,342 @@
package runtime
import (
"encoding/json"
"log/slog"
"strings"
"time"
"cs-agent/internal/models"
"cs-agent/internal/pkg/toolx"
svc "cs-agent/internal/services"
)
func newReplyRunLogService() *replyRunLogService {
return &replyRunLogService{}
}
type replyRunLogService struct{}
func (s *replyRunLogService) Write(startedAt time.Time, message models.Message, conversation models.Conversation, aiAgent models.AIAgent,
question string, runErr error, trace *aiReplyTraceData, summary *Summary) {
errorMessage := ""
if runErr != nil {
errorMessage = runErr.Error()
} else if summary != nil && strings.TrimSpace(summary.ErrorMessage) != "" {
errorMessage = strings.TrimSpace(summary.ErrorMessage)
}
traceData := buildAIReplyTraceData(trace)
plannedAction, plannedToolCode, planReason := buildRunLogPlan(summary)
logItem := &models.AgentRunLog{
ConversationID: conversation.ID,
MessageID: message.ID,
AIAgentID: aiAgent.ID,
AIConfigID: aiAgent.AIConfigID,
UserMessage: strings.TrimSpace(question),
PlannedAction: plannedAction,
PlannedSkillCode: strings.TrimSpace(summaryPlannedSkillCode(summary)),
PlannedSkillName: strings.TrimSpace(summaryPlannedSkillName(summary)),
SkillRouteTrace: strings.TrimSpace(summarySkillRouteTrace(summary)),
ToolSearchTrace: extractToolSearchTrace(summary),
GraphToolTrace: extractGraphToolTrace(summary),
GraphToolCode: firstGraphToolCode(summary),
HandoffReason: extractHandoffReason(summary),
PlannedToolCode: plannedToolCode,
PlanReason: planReason,
InterruptType: firstInterruptType(summary),
ResumeSource: runLogResumeSource(trace),
FinalAction: toRunLogFinalAction(summary),
FinalStatus: runLogFinalStatus(summary),
ReplyText: buildRunLogReplyText(summary),
ErrorMessage: errorMessage,
LatencyMs: time.Since(startedAt).Milliseconds(),
TraceData: traceData,
CreatedAt: time.Now(),
}
if err := svc.AgentRunLogService.Create(logItem); err != nil {
slog.Warn("create agent run log failed",
"message_id", message.ID,
"conversation_id", logItem.ConversationID,
"ai_agent_id", aiAgent.ID,
"error", err)
}
}
func buildAIReplyTraceData(trace *aiReplyTraceData) string {
if trace == nil {
return ""
}
data, err := json.Marshal(trace)
if err != nil {
return ""
}
return string(data)
}
func buildRunLogPlan(summary *Summary) (plannedAction, plannedToolCode, planReason string) {
if summary == nil {
return "", "", ""
}
if skillCode := strings.TrimSpace(summaryPlannedSkillCode(summary)); skillCode != "" {
reason := strings.TrimSpace(summary.PlanReason)
if reason == "" {
reason = "skill_selected"
}
return "skill", "", reason
}
if strings.TrimSpace(summary.Status) == "expired" {
return "interrupt", "", "pending interrupt checkpoint expired"
}
if summary.Interrupted {
if graphToolCode := firstGraphToolCode(summary); graphToolCode != "" {
reason := graphPlanReason(summary)
if reason == "" {
reason = "graph tool interrupted and is waiting for user confirmation"
}
return "graph", graphToolCode, reason
}
return "tool", summaryPrimaryToolCode(summary), "agent interrupted and is waiting for user confirmation"
}
if len(summary.InvokedToolCodes) > 0 {
if graphToolCode := firstGraphToolCode(summary); graphToolCode != "" {
reason := graphPlanReason(summary)
if reason == "" {
reason = "agent invoked graph tool"
}
return "graph", graphToolCode, reason
}
toolCode := summaryPrimaryToolCode(summary)
reason := "agent invoked MCP tool"
if toolCode != "" && toolCode != firstInvokedToolCode(summary) {
reason = "agent invoked dynamic tool via tool_search"
}
return "tool", toolCode, reason
}
if strings.TrimSpace(summary.ReplyText) != "" {
return "reply", "", "agent replied directly"
}
if strings.TrimSpace(summary.ErrorMessage) != "" {
return "error", "", "runtime execution failed"
}
return "fallback", "", "runtime produced empty reply"
}
func toRunLogFinalAction(summary *Summary) string {
if summary == nil {
return ""
}
if skillCode := strings.TrimSpace(summaryPlannedSkillCode(summary)); skillCode != "" && strings.TrimSpace(summary.ReplyText) != "" {
return "skill"
}
if graphToolCode := firstGraphToolCode(summary); graphToolCode != "" && strings.TrimSpace(summary.ReplyText) != "" {
return "graph"
}
switch strings.TrimSpace(summary.Status) {
case "completed":
return "reply"
case "fallback":
return "fallback"
case "error":
return "error"
case "interrupted":
return "interrupted"
case "expired":
return "expired"
default:
return strings.TrimSpace(summary.Status)
}
}
func buildRunLogReplyText(summary *Summary) string {
if summary == nil {
return ""
}
return strings.TrimSpace(summary.ReplyText)
}
func summaryPlannedSkillCode(summary *Summary) string {
if summary == nil {
return ""
}
return strings.TrimSpace(summary.PlannedSkillCode)
}
func summaryPlannedSkillName(summary *Summary) string {
if summary == nil {
return ""
}
return strings.TrimSpace(summary.PlannedSkillName)
}
func summarySkillRouteTrace(summary *Summary) string {
if summary == nil {
return ""
}
return strings.TrimSpace(summary.SkillRouteTrace)
}
func runLogResumeSource(trace *aiReplyTraceData) string {
if trace == nil {
return ""
}
return strings.TrimSpace(trace.ResumeSource)
}
func runLogFinalStatus(summary *Summary) string {
if summary == nil {
return ""
}
return strings.TrimSpace(summary.Status)
}
func summaryPrimaryToolCode(summary *Summary) string {
if summary == nil {
return ""
}
toolCode := firstInvokedToolCode(summary)
if toolCode != toolx.BuiltinToolSearch.Code {
return toolCode
}
if targetToolCode := firstToolSearchTargetToolCode(summary); targetToolCode != "" {
return targetToolCode
}
return toolCode
}
func extractToolSearchTrace(summary *Summary) string {
if summary == nil {
return ""
}
trace := parseRuntimeTraceData(summary.TraceData)
if len(trace.ToolSearch.Items) == 0 || len(trace.ToolSearch.Raw) == 0 {
return ""
}
return string(trace.ToolSearch.Raw)
}
func extractGraphToolTrace(summary *Summary) string {
if summary == nil {
return ""
}
trace := parseRuntimeTraceData(summary.TraceData)
if len(trace.GraphTools.Items) == 0 || len(trace.GraphTools.Raw) == 0 {
return ""
}
return string(trace.GraphTools.Raw)
}
func firstToolSearchTargetToolCode(summary *Summary) string {
trace := parseRuntimeTraceData(summary.TraceData)
for _, item := range trace.ToolSearch.Items {
toolCode := strings.TrimSpace(item.TargetToolCode)
if toolCode != "" {
return toolCode
}
if len(item.CandidateToolCodes) == 1 {
toolCode = strings.TrimSpace(item.CandidateToolCodes[0])
if toolCode != "" {
return toolCode
}
}
}
return ""
}
func firstGraphToolCode(summary *Summary) string {
trace := parseRuntimeTraceData(summary.TraceData)
for _, item := range trace.GraphTools.Items {
toolCode := strings.TrimSpace(item.ToolCode)
if toolCode != "" {
return toolCode
}
}
return ""
}
func extractHandoffReason(summary *Summary) string {
trace := parseRuntimeTraceData(summary.TraceData)
for _, item := range trace.GraphTools.Items {
if strings.TrimSpace(item.ToolCode) != toolx.GraphHandoffConversation.Code {
continue
}
if len(item.Arguments) == 0 {
return ""
}
reason, _ := item.Arguments["reason"].(string)
return strings.TrimSpace(reason)
}
return ""
}
func graphPlanReason(summary *Summary) string {
trace := parseRuntimeTraceData(summary.TraceData)
for _, item := range trace.GraphTools.Items {
toolCode := strings.TrimSpace(item.ToolCode)
switch toolCode {
case toolx.GraphTriageServiceRequest.Code:
recommendedAction := strings.TrimSpace(item.RecommendedAction)
if recommendedAction == "" {
return "graph tool triaged service request"
}
if item.TicketDraftReady {
return "graph tool triaged service request: " + recommendedAction + " with ready ticket draft"
}
return "graph tool triaged service request: " + recommendedAction
case toolx.GraphAnalyzeConversation.Code:
recommendedAction := strings.TrimSpace(item.RecommendedAction)
riskLevel := strings.TrimSpace(item.RiskLevel)
switch {
case recommendedAction != "" && riskLevel != "":
return "graph tool analyzed conversation: " + recommendedAction + " (" + riskLevel + " risk)"
case recommendedAction != "":
return "graph tool analyzed conversation: " + recommendedAction
case riskLevel != "":
return "graph tool analyzed conversation (" + riskLevel + " risk)"
default:
return "graph tool analyzed conversation"
}
}
}
return ""
}
type runtimeTraceProjection struct {
ToolSearch struct {
Raw json.RawMessage `json:"-"`
Items []struct {
TargetToolCode string `json:"targetToolCode"`
CandidateToolCodes []string `json:"candidateToolCodes"`
} `json:"items"`
} `json:"toolSearch"`
GraphTools struct {
Raw json.RawMessage `json:"-"`
Items []struct {
ToolCode string `json:"toolCode"`
Arguments map[string]any `json:"arguments"`
RecommendedAction string `json:"recommendedAction"`
RiskLevel string `json:"riskLevel"`
TicketDraftReady bool `json:"ticketDraftReady"`
} `json:"items"`
} `json:"graphTools"`
}
func parseRuntimeTraceData(raw string) runtimeTraceProjection {
raw = strings.TrimSpace(raw)
if raw == "" {
return runtimeTraceProjection{}
}
var payload map[string]json.RawMessage
if err := json.Unmarshal([]byte(raw), &payload); err != nil {
return runtimeTraceProjection{}
}
var trace runtimeTraceProjection
if toolSearchRaw, ok := payload["toolSearch"]; ok && len(toolSearchRaw) > 0 {
trace.ToolSearch.Raw = append(json.RawMessage(nil), toolSearchRaw...)
_ = json.Unmarshal(toolSearchRaw, &trace.ToolSearch)
}
if graphToolsRaw, ok := payload["graphTools"]; ok && len(graphToolsRaw) > 0 {
trace.GraphTools.Raw = append(json.RawMessage(nil), graphToolsRaw...)
_ = json.Unmarshal(graphToolsRaw, &trace.GraphTools)
}
return trace
}
+11 -543
View File
@@ -3,20 +3,13 @@ package runtime
import (
"context"
"encoding/json"
"fmt"
"log/slog"
"strings"
"time"
"cs-agent/internal/ai/runtime/graphs"
"cs-agent/internal/models"
"cs-agent/internal/pkg/dto"
"cs-agent/internal/pkg/enums"
"cs-agent/internal/pkg/toolx"
"cs-agent/internal/repositories"
svc "cs-agent/internal/services"
"github.com/mlogclub/simple/sqls"
)
var AIReplyService = newAIReplyService()
@@ -29,12 +22,18 @@ func newAIReplyService() *aiReplyService {
return &aiReplyService{
eligibility: newReplyEligibility(),
executor: newRuntimeReplyExecutor(),
interrupts: newReplyInterruptService(),
commit: newReplyCommitService(),
runlog: newReplyRunLogService(),
}
}
type aiReplyService struct {
eligibility *replyEligibility
executor *runtimeReplyExecutor
interrupts *replyInterruptService
commit *replyCommitService
runlog *replyRunLogService
}
type aiReplyTraceData struct {
@@ -95,10 +94,10 @@ func (s *aiReplyService) TriggerReply(ctx context.Context, conversation models.C
return nil
}
defer func() {
s.writeRunLog(startedAt, message, conversation, aiAgent, message.Content, retErr, trace, summary)
s.runlog.Write(startedAt, message, conversation, aiAgent, message.Content, retErr, trace, summary)
}()
if pendingInterrupt := svc.ConversationInterruptService.FindLatestPendingByConversationID(conversation.ID); pendingInterrupt != nil {
return s.resumePendingInterrupt(ctx, conversation, message, aiAgent, pendingInterrupt, trace, &summary)
return s.interrupts.ResumePendingInterrupt(ctx, s, conversation, message, aiAgent, pendingInterrupt, trace, &summary)
}
var err error
summary, err = s.executor.Run(ctx, conversation, message, aiAgent, trace)
@@ -106,14 +105,14 @@ func (s *aiReplyService) TriggerReply(ctx context.Context, conversation models.C
return err
}
if summary != nil && summary.Interrupted {
return s.handleInterruptedSummary(conversation, message, aiAgent, summary, trace)
return s.interrupts.HandleInterruptedSummary(s, conversation, message, aiAgent, summary, trace)
}
if summary != nil && strings.TrimSpace(summary.ReplyText) != "" {
replyMessage, err := s.sendAIReply(conversation, message, aiAgent, summary.ReplyText, trace, "ai_reply")
replyMessage, err := s.commit.SendAIReply(conversation, message, aiAgent, summary.ReplyText, trace, "ai_reply")
if err != nil {
return err
}
if err := s.incrementAIReplyRounds(conversation.ID, conversation.AIReplyRounds+1, aiAgent.Name); err != nil {
if err := s.commit.IncrementAIReplyRounds(conversation.ID, conversation.AIReplyRounds+1, aiAgent.Name); err != nil {
return err
}
trace.ReplySent = replyMessage != nil
@@ -121,365 +120,6 @@ func (s *aiReplyService) TriggerReply(ctx context.Context, conversation models.C
return nil
}
func (s *aiReplyService) resumePendingInterrupt(ctx context.Context, conversation models.Conversation, message models.Message, aiAgent models.AIAgent,
pendingInterrupt *models.ConversationInterrupt, trace *aiReplyTraceData, summaryRef **Summary) error {
if pendingInterrupt == nil {
return nil
}
summary, err := s.executor.ResumePendingInterrupt(ctx, conversation, message, aiAgent, pendingInterrupt, trace)
*summaryRef = summary
if err != nil {
if isCheckpointMissingError(err) {
summary = expiredInterruptSummary()
*summaryRef = summary
trace.Status = "interrupt_expired"
trace.FinalAction = "expired"
replyMessage, expireErr := s.sendAIReply(conversation, message, aiAgent, summary.ReplyText, trace, "ai_interrupt_expired")
if expireErr != nil {
return expireErr
}
if err := s.incrementAIReplyRounds(conversation.ID, conversation.AIReplyRounds+1, aiAgent.Name); err != nil {
return err
}
lastResumeMessageID := int64(0)
if replyMessage != nil {
lastResumeMessageID = replyMessage.ID
}
if expireMarkErr := svc.ConversationInterruptService.MarkExpired(pendingInterrupt.ID, lastResumeMessageID); expireMarkErr != nil {
return expireMarkErr
}
return nil
}
return err
}
if summary != nil && summary.Interrupted {
return s.handleInterruptedResume(conversation, message, aiAgent, pendingInterrupt, summary, trace)
}
if summary != nil && strings.TrimSpace(summary.ReplyText) != "" {
replyMessage, err := s.sendAIReply(conversation, message, aiAgent, summary.ReplyText, trace, "ai_resume")
if err != nil {
return err
}
if err := s.incrementAIReplyRounds(conversation.ID, conversation.AIReplyRounds+1, aiAgent.Name); err != nil {
return err
}
replyMessageID := int64(0)
if replyMessage != nil {
replyMessageID = replyMessage.ID
}
if graphs.IsCancellationReply(summary.ReplyText) {
return svc.ConversationInterruptService.MarkCancelled(pendingInterrupt.ID, replyMessageID)
}
return svc.ConversationInterruptService.MarkResolved(pendingInterrupt.ID, replyMessageID)
}
return svc.ConversationInterruptService.MarkResolved(pendingInterrupt.ID, 0)
}
func (s *aiReplyService) handleInterruptedSummary(conversation models.Conversation, message models.Message, aiAgent models.AIAgent,
summary *Summary, trace *aiReplyTraceData) error {
pending := buildConversationInterrupt(conversation, message, aiAgent, summary)
if err := svc.ConversationInterruptService.CreateOrUpdatePending(pending); err != nil {
return err
}
pending = svc.ConversationInterruptService.GetByCheckPointID(summary.CheckPointID)
replyText := resolveInterruptPrompt(summary)
replyMessage, err := s.sendAIReply(conversation, message, aiAgent, replyText, trace, "ai_interrupt")
if err != nil {
return err
}
if err := s.incrementAIReplyRounds(conversation.ID, conversation.AIReplyRounds+1, aiAgent.Name); err != nil {
return err
}
if replyMessage != nil && pending != nil {
return svc.ConversationInterruptService.MarkPendingAgain(pending.ID, pending.InterruptID, replyText, replyMessage.ID)
}
return nil
}
func (s *aiReplyService) handleInterruptedResume(conversation models.Conversation, message models.Message, aiAgent models.AIAgent,
pendingInterrupt *models.ConversationInterrupt, summary *Summary, trace *aiReplyTraceData) error {
if pendingInterrupt == nil {
return nil
}
replyText := resolveInterruptPrompt(summary)
replyMessage, err := s.sendAIReply(conversation, message, aiAgent, replyText, trace, "ai_interrupt_resume")
if err != nil {
return err
}
if err := s.incrementAIReplyRounds(conversation.ID, conversation.AIReplyRounds+1, aiAgent.Name); err != nil {
return err
}
if replyMessage != nil {
return svc.ConversationInterruptService.MarkPendingAgain(pendingInterrupt.ID, firstInterruptID(summary), replyText, replyMessage.ID)
}
return nil
}
func (s *aiReplyService) sendAIReply(conversation models.Conversation, message models.Message, aiAgent models.AIAgent,
replyText string, trace *aiReplyTraceData, clientPrefix string) (*models.Message, error) {
replyText = strings.TrimSpace(replyText)
if replyText == "" {
return nil, nil
}
commitStartedAt := time.Now()
replyMessage, err := svc.MessageService.SendAIMessage(conversation.ID, aiAgent.ID,
fmt.Sprintf("%s_%d", strings.TrimSpace(clientPrefix), message.ID), enums.IMMessageTypeText, replyText, "", s.buildAIPrincipal(aiAgent))
if trace != nil {
trace.CommitMs = time.Since(commitStartedAt).Milliseconds()
trace.ReplySent = err == nil && replyMessage != nil
if replyMessage != nil {
trace.ReplyMessageID = replyMessage.ID
}
}
return replyMessage, err
}
func (s *aiReplyService) writeRunLog(startedAt time.Time, message models.Message, conversation models.Conversation, aiAgent models.AIAgent,
question string, runErr error, trace *aiReplyTraceData, summary *Summary) {
errorMessage := ""
if runErr != nil {
errorMessage = runErr.Error()
} else if summary != nil && strings.TrimSpace(summary.ErrorMessage) != "" {
errorMessage = strings.TrimSpace(summary.ErrorMessage)
}
traceData := buildAIReplyTraceData(trace)
plannedAction, plannedToolCode, planReason := buildRunLogPlan(summary)
logItem := &models.AgentRunLog{
ConversationID: conversation.ID,
MessageID: message.ID,
AIAgentID: aiAgent.ID,
AIConfigID: aiAgent.AIConfigID,
UserMessage: strings.TrimSpace(question),
PlannedAction: plannedAction,
PlannedSkillCode: strings.TrimSpace(summaryPlannedSkillCode(summary)),
PlannedSkillName: strings.TrimSpace(summaryPlannedSkillName(summary)),
SkillRouteTrace: strings.TrimSpace(summarySkillRouteTrace(summary)),
ToolSearchTrace: extractToolSearchTrace(summary),
GraphToolTrace: extractGraphToolTrace(summary),
GraphToolCode: firstGraphToolCode(summary),
HandoffReason: extractHandoffReason(summary),
PlannedToolCode: plannedToolCode,
PlanReason: planReason,
InterruptType: firstInterruptType(summary),
ResumeSource: runLogResumeSource(trace),
FinalAction: toRunLogFinalAction(summary),
FinalStatus: runLogFinalStatus(summary),
ReplyText: buildRunLogReplyText(summary),
ErrorMessage: errorMessage,
LatencyMs: time.Since(startedAt).Milliseconds(),
TraceData: traceData,
CreatedAt: time.Now(),
}
if err := svc.AgentRunLogService.Create(logItem); err != nil {
slog.Warn("create agent run log failed",
"message_id", message.ID,
"conversation_id", logItem.ConversationID,
"ai_agent_id", aiAgent.ID,
"error", err)
}
}
func buildAIReplyTraceData(trace *aiReplyTraceData) string {
if trace == nil {
return ""
}
data, err := json.Marshal(trace)
if err != nil {
return ""
}
return string(data)
}
func buildRunLogPlan(summary *Summary) (plannedAction, plannedToolCode, planReason string) {
if summary == nil {
return "", "", ""
}
if skillCode := strings.TrimSpace(summaryPlannedSkillCode(summary)); skillCode != "" {
reason := strings.TrimSpace(summary.PlanReason)
if reason == "" {
reason = "skill_selected"
}
return "skill", "", reason
}
if strings.TrimSpace(summary.Status) == "expired" {
return "interrupt", "", "pending interrupt checkpoint expired"
}
if summary.Interrupted {
if graphToolCode := firstGraphToolCode(summary); graphToolCode != "" {
reason := graphPlanReason(summary)
if reason == "" {
reason = "graph tool interrupted and is waiting for user confirmation"
}
return "graph", graphToolCode, reason
}
return "tool", summaryPrimaryToolCode(summary), "agent interrupted and is waiting for user confirmation"
}
if len(summary.InvokedToolCodes) > 0 {
if graphToolCode := firstGraphToolCode(summary); graphToolCode != "" {
reason := graphPlanReason(summary)
if reason == "" {
reason = "agent invoked graph tool"
}
return "graph", graphToolCode, reason
}
toolCode := summaryPrimaryToolCode(summary)
reason := "agent invoked MCP tool"
if toolCode != "" && toolCode != firstInvokedToolCode(summary) {
reason = "agent invoked dynamic tool via tool_search"
}
return "tool", toolCode, reason
}
if strings.TrimSpace(summary.ReplyText) != "" {
return "reply", "", "agent replied directly"
}
if strings.TrimSpace(summary.ErrorMessage) != "" {
return "error", "", "runtime execution failed"
}
return "fallback", "", "runtime produced empty reply"
}
func toRunLogFinalAction(summary *Summary) string {
if summary == nil {
return ""
}
if skillCode := strings.TrimSpace(summaryPlannedSkillCode(summary)); skillCode != "" && strings.TrimSpace(summary.ReplyText) != "" {
return "skill"
}
if graphToolCode := firstGraphToolCode(summary); graphToolCode != "" && strings.TrimSpace(summary.ReplyText) != "" {
return "graph"
}
switch strings.TrimSpace(summary.Status) {
case "completed":
return "reply"
case "fallback":
return "fallback"
case "error":
return "error"
case "interrupted":
return "interrupted"
case "expired":
return "expired"
default:
return strings.TrimSpace(summary.Status)
}
}
func buildRunLogReplyText(summary *Summary) string {
if summary == nil {
return ""
}
return strings.TrimSpace(summary.ReplyText)
}
func summaryPlannedSkillCode(summary *Summary) string {
if summary == nil {
return ""
}
return strings.TrimSpace(summary.PlannedSkillCode)
}
func summaryPlannedSkillName(summary *Summary) string {
if summary == nil {
return ""
}
return strings.TrimSpace(summary.PlannedSkillName)
}
func summarySkillRouteTrace(summary *Summary) string {
if summary == nil {
return ""
}
return strings.TrimSpace(summary.SkillRouteTrace)
}
func runLogResumeSource(trace *aiReplyTraceData) string {
if trace == nil {
return ""
}
return strings.TrimSpace(trace.ResumeSource)
}
func runLogFinalStatus(summary *Summary) string {
if summary == nil {
return ""
}
return strings.TrimSpace(summary.Status)
}
func (s *aiReplyService) incrementAIReplyRounds(conversationID int64, nextRounds int, aiAgentName string) error {
return repositories.ConversationRepository.Updates(sqls.DB(), conversationID, map[string]any{
"ai_reply_rounds": nextRounds,
"update_user_id": 0,
"update_user_name": strings.TrimSpace(aiAgentName),
"updated_at": time.Now(),
})
}
func buildConversationInterrupt(conversation models.Conversation, message models.Message, aiAgent models.AIAgent, summary *Summary) *models.ConversationInterrupt {
if summary == nil {
return nil
}
now := time.Now()
item := svc.ConversationInterruptService.GetByCheckPointID(summary.CheckPointID)
if item == nil {
item = &models.ConversationInterrupt{
CheckPointID: summary.CheckPointID,
CreatedAt: now,
}
}
item.ConversationID = conversation.ID
item.AIAgentID = aiAgent.ID
item.SourceMessageID = message.ID
item.InterruptID = firstInterruptID(summary)
item.InterruptType = firstInterruptType(summary)
item.Status = "pending"
item.PromptText = resolveInterruptPrompt(summary)
item.UpdatedAt = now
return item
}
func resolveInterruptPrompt(summary *Summary) string {
if summary == nil || len(summary.Interrupts) == 0 {
return "请继续补充信息后再试。"
}
if prompt := extractInterruptMessage(summary.Interrupts[0].InfoPreview); prompt != "" {
return prompt
}
if prompt := strings.TrimSpace(summary.Interrupts[0].InfoPreview); prompt != "" {
return prompt
}
return "请继续补充信息后再试。"
}
func extractInterruptMessage(infoPreview string) string {
infoPreview = strings.TrimSpace(infoPreview)
if infoPreview == "" {
return ""
}
payload := make(map[string]any)
if err := json.Unmarshal([]byte(infoPreview), &payload); err != nil {
return ""
}
if message, ok := payload["message"].(string); ok {
return strings.TrimSpace(message)
}
return ""
}
func firstInterruptID(summary *Summary) string {
if summary == nil || len(summary.Interrupts) == 0 {
return ""
}
return strings.TrimSpace(summary.Interrupts[0].ID)
}
func firstInterruptType(summary *Summary) string {
if summary == nil || len(summary.Interrupts) == 0 {
return ""
}
return strings.TrimSpace(summary.Interrupts[0].Type)
}
func firstInvokedToolCode(summary *Summary) string {
if summary == nil {
return ""
@@ -489,175 +129,3 @@ func firstInvokedToolCode(summary *Summary) string {
}
return ""
}
func summaryPrimaryToolCode(summary *Summary) string {
if summary == nil {
return ""
}
toolCode := firstInvokedToolCode(summary)
if toolCode != toolx.BuiltinToolSearch.Code {
return toolCode
}
if targetToolCode := firstToolSearchTargetToolCode(summary); targetToolCode != "" {
return targetToolCode
}
return toolCode
}
func extractToolSearchTrace(summary *Summary) string {
if summary == nil {
return ""
}
trace := parseRuntimeTraceData(summary.TraceData)
if len(trace.ToolSearch.Items) == 0 || len(trace.ToolSearch.Raw) == 0 {
return ""
}
return string(trace.ToolSearch.Raw)
}
func extractGraphToolTrace(summary *Summary) string {
if summary == nil {
return ""
}
trace := parseRuntimeTraceData(summary.TraceData)
if len(trace.GraphTools.Items) == 0 || len(trace.GraphTools.Raw) == 0 {
return ""
}
return string(trace.GraphTools.Raw)
}
func firstToolSearchTargetToolCode(summary *Summary) string {
trace := parseRuntimeTraceData(summary.TraceData)
for _, item := range trace.ToolSearch.Items {
toolCode := strings.TrimSpace(item.TargetToolCode)
if toolCode != "" {
return toolCode
}
if len(item.CandidateToolCodes) == 1 {
toolCode = strings.TrimSpace(item.CandidateToolCodes[0])
if toolCode != "" {
return toolCode
}
}
}
return ""
}
func firstGraphToolCode(summary *Summary) string {
trace := parseRuntimeTraceData(summary.TraceData)
for _, item := range trace.GraphTools.Items {
toolCode := strings.TrimSpace(item.ToolCode)
if toolCode != "" {
return toolCode
}
}
return ""
}
func extractHandoffReason(summary *Summary) string {
trace := parseRuntimeTraceData(summary.TraceData)
for _, item := range trace.GraphTools.Items {
if strings.TrimSpace(item.ToolCode) != toolx.GraphHandoffConversation.Code {
continue
}
if len(item.Arguments) == 0 {
return ""
}
reason, _ := item.Arguments["reason"].(string)
return strings.TrimSpace(reason)
}
return ""
}
func graphPlanReason(summary *Summary) string {
trace := parseRuntimeTraceData(summary.TraceData)
for _, item := range trace.GraphTools.Items {
toolCode := strings.TrimSpace(item.ToolCode)
switch toolCode {
case toolx.GraphTriageServiceRequest.Code:
recommendedAction := strings.TrimSpace(item.RecommendedAction)
if recommendedAction == "" {
return "graph tool triaged service request"
}
if item.TicketDraftReady {
return "graph tool triaged service request: " + recommendedAction + " with ready ticket draft"
}
return "graph tool triaged service request: " + recommendedAction
case toolx.GraphAnalyzeConversation.Code:
recommendedAction := strings.TrimSpace(item.RecommendedAction)
riskLevel := strings.TrimSpace(item.RiskLevel)
switch {
case recommendedAction != "" && riskLevel != "":
return "graph tool analyzed conversation: " + recommendedAction + " (" + riskLevel + " risk)"
case recommendedAction != "":
return "graph tool analyzed conversation: " + recommendedAction
case riskLevel != "":
return "graph tool analyzed conversation (" + riskLevel + " risk)"
default:
return "graph tool analyzed conversation"
}
}
}
return ""
}
type runtimeTraceProjection struct {
ToolSearch struct {
Raw json.RawMessage `json:"-"`
Items []struct {
TargetToolCode string `json:"targetToolCode"`
CandidateToolCodes []string `json:"candidateToolCodes"`
} `json:"items"`
} `json:"toolSearch"`
GraphTools struct {
Raw json.RawMessage `json:"-"`
Items []struct {
ToolCode string `json:"toolCode"`
Arguments map[string]any `json:"arguments"`
RecommendedAction string `json:"recommendedAction"`
RiskLevel string `json:"riskLevel"`
TicketDraftReady bool `json:"ticketDraftReady"`
} `json:"items"`
} `json:"graphTools"`
}
func parseRuntimeTraceData(raw string) runtimeTraceProjection {
raw = strings.TrimSpace(raw)
if raw == "" {
return runtimeTraceProjection{}
}
var payload map[string]json.RawMessage
if err := json.Unmarshal([]byte(raw), &payload); err != nil {
return runtimeTraceProjection{}
}
var trace runtimeTraceProjection
if toolSearchRaw, ok := payload["toolSearch"]; ok && len(toolSearchRaw) > 0 {
trace.ToolSearch.Raw = append(json.RawMessage(nil), toolSearchRaw...)
_ = json.Unmarshal(toolSearchRaw, &trace.ToolSearch)
}
if graphToolsRaw, ok := payload["graphTools"]; ok && len(graphToolsRaw) > 0 {
trace.GraphTools.Raw = append(json.RawMessage(nil), graphToolsRaw...)
_ = json.Unmarshal(graphToolsRaw, &trace.GraphTools)
}
return trace
}
func isCheckpointMissingError(err error) bool {
if err == nil {
return false
}
message := strings.ToLower(strings.TrimSpace(err.Error()))
return strings.Contains(message, "failed to load from checkpoint") && strings.Contains(message, "not exist")
}
func (s *aiReplyService) buildAIPrincipal(aiAgent models.AIAgent) *dto.AuthPrincipal {
username := "AI"
if strings.TrimSpace(aiAgent.Name) != "" {
username = aiAgent.Name
}
return &dto.AuthPrincipal{
UserID: 0,
Username: username,
Nickname: username,
}
}
+161
View File
@@ -0,0 +1,161 @@
package runtime
import (
"testing"
"time"
"cs-agent/internal/models"
"cs-agent/internal/pkg/enums"
"cs-agent/internal/pkg/toolx"
)
func TestReplyEligibilityCanReply(t *testing.T) {
eligibility := newReplyEligibility()
conversation := newConversationFixture()
message := newCustomerMessageFixture("hello")
aiAgent := newAIAgentFixture()
if !eligibility.CanReply(conversation, message, aiAgent) {
t.Fatalf("expected customer message to be replyable")
}
message.SenderType = enums.IMSenderTypeAgent
if eligibility.CanReply(conversation, message, aiAgent) {
t.Fatalf("expected non-customer message to be rejected")
}
message = newCustomerMessageFixture("hello")
conversation.HandoffAt = ptrTime(time.Now())
if eligibility.CanReply(conversation, message, aiAgent) {
t.Fatalf("expected handed-off conversation to be rejected")
}
conversation = newConversationFixture()
conversation.CurrentAssigneeID = 1
if eligibility.CanReply(conversation, message, aiAgent) {
t.Fatalf("expected assigned conversation to be rejected")
}
conversation = newConversationFixture()
aiAgent.ServiceMode = enums.IMConversationServiceModeHumanOnly
if eligibility.CanReply(conversation, message, aiAgent) {
t.Fatalf("expected human-only agent to be rejected")
}
aiAgent = newAIAgentFixture()
message.Content = " "
if eligibility.CanReply(conversation, message, aiAgent) {
t.Fatalf("expected blank message to be rejected")
}
}
func TestResolveReplyTimeout(t *testing.T) {
service := newAIReplyService()
aiAgent := newAIAgentFixture()
if got := service.resolveReplyTimeout(aiAgent); got != 180*time.Second {
t.Fatalf("expected default timeout, got %v", got)
}
aiAgent.ReplyTimeoutSeconds = 30
if got := service.resolveReplyTimeout(aiAgent); got != 30*time.Second {
t.Fatalf("expected exact timeout, got %v", got)
}
aiAgent.ReplyTimeoutSeconds = 999
if got := service.resolveReplyTimeout(aiAgent); got != 600*time.Second {
t.Fatalf("expected clamped timeout, got %v", got)
}
}
func TestBuildRunLogPlan(t *testing.T) {
summary := &Summary{
PlannedSkillCode: "faq_router",
PlanReason: "manual",
}
action, toolCode, reason := buildRunLogPlan(summary)
if action != "skill" || toolCode != "" || reason != "manual" {
t.Fatalf("unexpected skill plan result: action=%q toolCode=%q reason=%q", action, toolCode, reason)
}
summary = &Summary{
Interrupted: true,
TraceData: `{
"graphTools": {
"items": [
{
"toolCode": "` + toolx.GraphTriageServiceRequest.Code + `",
"recommendedAction": "create_ticket",
"ticketDraftReady": true
}
]
}
}`,
}
action, toolCode, reason = buildRunLogPlan(summary)
if action != "graph" || toolCode != toolx.GraphTriageServiceRequest.Code || reason == "" {
t.Fatalf("unexpected graph interrupt result: action=%q toolCode=%q reason=%q", action, toolCode, reason)
}
summary = &Summary{
InvokedToolCodes: []string{toolx.BuiltinToolSearch.Code},
TraceData: `{
"toolSearch": {
"items": [
{
"targetToolCode": "mcp/test/search"
}
]
}
}`,
}
action, toolCode, reason = buildRunLogPlan(summary)
if action != "tool" || toolCode != "mcp/test/search" || reason != "agent invoked dynamic tool via tool_search" {
t.Fatalf("unexpected dynamic tool result: action=%q toolCode=%q reason=%q", action, toolCode, reason)
}
summary = &Summary{ReplyText: "done"}
action, toolCode, reason = buildRunLogPlan(summary)
if action != "reply" || toolCode != "" || reason != "agent replied directly" {
t.Fatalf("unexpected reply result: action=%q toolCode=%q reason=%q", action, toolCode, reason)
}
}
func TestResolveInterruptPrompt(t *testing.T) {
summary := &Summary{
Interrupts: []InterruptContextSummary{
{
ID: "interrupt-1",
Type: "question",
InfoPreview: `{"message":"请补充订单号"}`,
},
},
}
if got := resolveInterruptPrompt(summary); got != "请补充订单号" {
t.Fatalf("unexpected interrupt prompt: %q", got)
}
summary.Interrupts[0].InfoPreview = "直接补充手机号"
if got := resolveInterruptPrompt(summary); got != "直接补充手机号" {
t.Fatalf("unexpected raw interrupt prompt: %q", got)
}
}
func newConversationFixture() models.Conversation {
return models.Conversation{}
}
func newCustomerMessageFixture(content string) models.Message {
return models.Message{
SenderType: enums.IMSenderTypeCustomer,
Content: content,
}
}
func newAIAgentFixture() models.AIAgent {
return models.AIAgent{}
}
func ptrTime(v time.Time) *time.Time {
return &v
}
-107
View File
@@ -2,12 +2,8 @@ package skills
import (
"context"
"encoding/json"
"fmt"
"strings"
"time"
"cs-agent/internal/ai"
"cs-agent/internal/models"
"cs-agent/internal/pkg/enums"
"cs-agent/internal/pkg/errorsx"
@@ -106,106 +102,3 @@ func loadCandidateSkills(aiAgent *models.AIAgent) []models.SkillDefinition {
}
return ret
}
func routeSkillWithLLM(ctx context.Context, aiConfig *models.AIConfig, userMessage string, candidates []models.SkillDefinition) (*models.SkillDefinition, *RouteTrace, error) {
trace := &RouteTrace{Status: "started"}
if aiConfig == nil {
trace.Status = "config_error"
trace.Error = "ai config is nil"
return nil, trace, errorsx.InvalidParam("Skill 路由依赖的 AI 配置不可用")
}
if len(candidates) == 0 {
trace.Status = "no_candidate"
return nil, trace, nil
}
userMessage = strings.TrimSpace(userMessage)
if userMessage == "" {
trace.Status = "empty_user_message"
return nil, trace, nil
}
systemPrompt := "你是客服技能路由器。你只能在候选 Skill 中选择一个最合适的 skillCode,或者返回 NONE。只有当用户问题与 Skill 的职责边界明确匹配时才选择;如果不明确、信息不足、多个 Skill 都不够确定,就返回 NONE。输出只能是 skillCode 或 NONE,不能输出其他内容。"
userPrompt := buildSkillRoutePrompt(userMessage, candidates)
startedAt := time.Now()
result, err := ai.LLM.ChatWithConfig(ctx, aiConfig, systemPrompt, userPrompt)
trace.LatencyMs = time.Since(startedAt).Milliseconds()
if err != nil {
trace.Status = "route_error"
trace.Error = err.Error()
return nil, trace, err
}
decision := normalizeRouteDecision(result.Content)
trace.RawDecision = strings.TrimSpace(result.Content)
if decision == "" || decision == "NONE" {
trace.Status = "not_matched"
return nil, trace, nil
}
for _, item := range candidates {
if strings.EqualFold(item.Code, decision) {
trace.Status = "llm_selected"
trace.SelectedSkillCode = item.Code
return &item, trace, nil
}
}
trace.Status = "invalid_decision"
trace.Error = fmt.Sprintf("invalid route decision: %s", decision)
return nil, trace, nil
}
func buildSkillRoutePrompt(userMessage string, candidates []models.SkillDefinition) string {
lines := make([]string, 0, len(candidates)+4)
lines = append(lines, "用户问题:")
lines = append(lines, strings.TrimSpace(userMessage))
lines = append(lines, "")
lines = append(lines, "候选 Skills")
for _, item := range candidates {
line := fmt.Sprintf("- skillCode=%s; name=%s; description=%s", strings.TrimSpace(item.Code), strings.TrimSpace(item.Name), strings.TrimSpace(item.Description))
if examples := parseSkillExamples(item.Examples); len(examples) > 0 {
line += "; examples=" + strings.Join(examples, " | ")
}
lines = append(lines, line)
}
lines = append(lines, "")
lines = append(lines, "请只输出一个 skillCode 或 NONE。")
return strings.Join(lines, "\n")
}
func parseSkillExamples(raw string) []string {
raw = strings.TrimSpace(raw)
if raw == "" {
return nil
}
var items []string
if err := json.Unmarshal([]byte(raw), &items); err != nil {
return nil
}
ret := make([]string, 0, len(items))
for _, item := range items {
item = strings.TrimSpace(item)
if item == "" {
continue
}
ret = append(ret, item)
if len(ret) >= 3 {
break
}
}
return ret
}
func normalizeRouteDecision(raw string) string {
raw = strings.TrimSpace(raw)
if raw == "" {
return ""
}
raw = strings.Trim(raw, "`")
raw = strings.TrimSpace(raw)
if idx := strings.Index(raw, "\n"); idx >= 0 {
raw = raw[:idx]
}
raw = strings.TrimSpace(raw)
raw = strings.Trim(raw, "\"'")
if strings.EqualFold(raw, "NONE") {
return "NONE"
}
return raw
}
+116
View File
@@ -0,0 +1,116 @@
package skills
import (
"context"
"encoding/json"
"fmt"
"strings"
"time"
"cs-agent/internal/ai"
"cs-agent/internal/models"
"cs-agent/internal/pkg/errorsx"
)
func routeSkillWithLLM(ctx context.Context, aiConfig *models.AIConfig, userMessage string, candidates []models.SkillDefinition) (*models.SkillDefinition, *RouteTrace, error) {
trace := &RouteTrace{Status: "started"}
if aiConfig == nil {
trace.Status = "config_error"
trace.Error = "ai config is nil"
return nil, trace, errorsx.InvalidParam("Skill 路由依赖的 AI 配置不可用")
}
if len(candidates) == 0 {
trace.Status = "no_candidate"
return nil, trace, nil
}
userMessage = strings.TrimSpace(userMessage)
if userMessage == "" {
trace.Status = "empty_user_message"
return nil, trace, nil
}
systemPrompt := "你是客服技能路由器。你只能在候选 Skill 中选择一个最合适的 skillCode,或者返回 NONE。只有当用户问题与 Skill 的职责边界明确匹配时才选择;如果不明确、信息不足、多个 Skill 都不够确定,就返回 NONE。输出只能是 skillCode 或 NONE,不能输出其他内容。"
userPrompt := buildSkillRoutePrompt(userMessage, candidates)
startedAt := time.Now()
result, err := ai.LLM.ChatWithConfig(ctx, aiConfig, systemPrompt, userPrompt)
trace.LatencyMs = time.Since(startedAt).Milliseconds()
if err != nil {
trace.Status = "route_error"
trace.Error = err.Error()
return nil, trace, err
}
decision := normalizeRouteDecision(result.Content)
trace.RawDecision = strings.TrimSpace(result.Content)
if decision == "" || decision == "NONE" {
trace.Status = "not_matched"
return nil, trace, nil
}
for _, item := range candidates {
if strings.EqualFold(item.Code, decision) {
trace.Status = "llm_selected"
trace.SelectedSkillCode = item.Code
return &item, trace, nil
}
}
trace.Status = "invalid_decision"
trace.Error = fmt.Sprintf("invalid route decision: %s", decision)
return nil, trace, nil
}
func buildSkillRoutePrompt(userMessage string, candidates []models.SkillDefinition) string {
lines := make([]string, 0, len(candidates)+4)
lines = append(lines, "用户问题:")
lines = append(lines, strings.TrimSpace(userMessage))
lines = append(lines, "")
lines = append(lines, "候选 Skills")
for _, item := range candidates {
line := fmt.Sprintf("- skillCode=%s; name=%s; description=%s", strings.TrimSpace(item.Code), strings.TrimSpace(item.Name), strings.TrimSpace(item.Description))
if examples := parseSkillExamples(item.Examples); len(examples) > 0 {
line += "; examples=" + strings.Join(examples, " | ")
}
lines = append(lines, line)
}
lines = append(lines, "")
lines = append(lines, "请只输出一个 skillCode 或 NONE。")
return strings.Join(lines, "\n")
}
func parseSkillExamples(raw string) []string {
raw = strings.TrimSpace(raw)
if raw == "" {
return nil
}
var items []string
if err := json.Unmarshal([]byte(raw), &items); err != nil {
return nil
}
ret := make([]string, 0, len(items))
for _, item := range items {
item = strings.TrimSpace(item)
if item == "" {
continue
}
ret = append(ret, item)
if len(ret) >= 3 {
break
}
}
return ret
}
func normalizeRouteDecision(raw string) string {
raw = strings.TrimSpace(raw)
if raw == "" {
return ""
}
if idx := strings.Index(raw, "\n"); idx >= 0 {
raw = raw[:idx]
}
raw = strings.TrimSpace(raw)
raw = strings.Trim(raw, "`")
raw = strings.TrimSpace(raw)
raw = strings.Trim(raw, "\"'")
if strings.EqualFold(raw, "NONE") {
return "NONE"
}
return raw
}
+48
View File
@@ -0,0 +1,48 @@
package skills
import (
"strings"
"testing"
"cs-agent/internal/models"
)
func TestParseSkillExamples(t *testing.T) {
examples := parseSkillExamples(`[" 退款进度 ","","发票补开","修改收货地址","多余示例"]`)
if len(examples) != 3 {
t.Fatalf("expected 3 examples, got %d", len(examples))
}
if examples[0] != "退款进度" || examples[1] != "发票补开" || examples[2] != "修改收货地址" {
t.Fatalf("unexpected examples: %#v", examples)
}
}
func TestNormalizeRouteDecision(t *testing.T) {
if got := normalizeRouteDecision("```refund_skill```\n补充说明"); got != "refund_skill" {
t.Fatalf("unexpected normalized decision: %q", got)
}
if got := normalizeRouteDecision(" none "); got != "NONE" {
t.Fatalf("expected NONE, got %q", got)
}
}
func TestBuildSkillRoutePrompt(t *testing.T) {
prompt := buildSkillRoutePrompt("我要申请退款", []models.SkillDefinition{
{
Code: "refund_skill",
Name: "退款处理",
Description: "负责退款和退货相关问题",
Examples: `["退款进度","退货运费"]`,
},
})
if !strings.Contains(prompt, "skillCode=refund_skill") {
t.Fatalf("expected prompt to include skill code, got %q", prompt)
}
if !strings.Contains(prompt, "examples=退款进度 | 退货运费") {
t.Fatalf("expected prompt to include examples, got %q", prompt)
}
if !strings.Contains(prompt, "请只输出一个 skillCode 或 NONE。") {
t.Fatalf("expected prompt to include output constraint, got %q", prompt)
}
}