From e5448e00b3b686e1ed7e9cbc280699e685e3629f Mon Sep 17 00:00:00 2001 From: mlogclub Date: Mon, 13 Apr 2026 18:49:05 +0800 Subject: [PATCH] Refactor execution logic into a new executor package - Moved execution-related logic from the engine service to a new executor service. - Created types for RunInput, ResumeInput, RunResult, and InterruptContextSummary in the executor package. - Updated the engine service to utilize the new executor service for executing runs and resumes. - Removed redundant code and improved organization by separating concerns between engine and executor. --- internal/ai/runtime/app/service.go | 12 +- .../ai/runtime/internal/engine/service.go | 549 +----------------- internal/ai/runtime/internal/engine/types.go | 60 +- .../ai/runtime/internal/executor/service.go | 541 +++++++++++++++++ .../ai/runtime/internal/executor/types.go | 57 ++ 5 files changed, 615 insertions(+), 604 deletions(-) create mode 100644 internal/ai/runtime/internal/executor/service.go create mode 100644 internal/ai/runtime/internal/executor/types.go diff --git a/internal/ai/runtime/app/service.go b/internal/ai/runtime/app/service.go index 9ebf644..639ebaf 100644 --- a/internal/ai/runtime/app/service.go +++ b/internal/ai/runtime/app/service.go @@ -5,7 +5,7 @@ import ( "encoding/json" "strings" - "cs-agent/internal/ai/runtime/internal/engine" + "cs-agent/internal/ai/runtime/internal/executor" "cs-agent/internal/ai/runtime/registry" "cs-agent/internal/ai/runtime/tools" "cs-agent/internal/ai/skills" @@ -14,13 +14,13 @@ import ( ) type Service struct { - runtime *engine.Service + runtime *executor.Service registry *registry.Registry } func NewService() *Service { return &Service{ - runtime: engine.NewService(), + runtime: executor.NewService(), registry: registry.NewRegistry( tools.NewTriageServiceRequestTool(), tools.NewAnalyzeConversationTool(), @@ -42,7 +42,7 @@ func (s *Service) Run(ctx context.Context, req Request) (*Summary, error) { if err := s.prepareToolsForRun(&req); err != nil { return nil, err } - summary, err := s.runtime.ExecuteRun(ctx, engine.RunInput{ + summary, err := s.runtime.ExecuteRun(ctx, executor.RunInput{ Conversation: req.Conversation, UserMessage: req.UserMessage, AIAgent: req.AIAgent, @@ -71,7 +71,7 @@ func (s *Service) Resume(ctx context.Context, req ResumeRequest) (*Summary, erro if err := s.prepareToolsForResume(&req); err != nil { return nil, err } - summary, err := s.runtime.ExecuteResume(ctx, engine.ResumeInput{ + summary, err := s.runtime.ExecuteResume(ctx, executor.ResumeInput{ Conversation: req.Conversation, AIAgent: req.AIAgent, AIConfig: req.AIConfig, @@ -120,7 +120,7 @@ func (s *Service) prepareToolsForResume(req *ResumeRequest) error { return nil } -func toSummary(summary *engine.Summary) *Summary { +func toSummary(summary *executor.RunResult) *Summary { if summary == nil { return nil } diff --git a/internal/ai/runtime/internal/engine/service.go b/internal/ai/runtime/internal/engine/service.go index e1c1c20..c76ab1b 100644 --- a/internal/ai/runtime/internal/engine/service.go +++ b/internal/ai/runtime/internal/engine/service.go @@ -2,34 +2,17 @@ package engine import ( "context" - "encoding/json" - "fmt" - "strings" - "cs-agent/internal/ai/runtime/internal/impl/adapter" - "cs-agent/internal/ai/runtime/internal/impl/callbacks" - "cs-agent/internal/ai/runtime/internal/impl/factory" - "cs-agent/internal/ai/runtime/internal/impl/retrievers" - "cs-agent/internal/ai/runtime/registry" - "cs-agent/internal/models" - "cs-agent/internal/pkg/toolx" - "cs-agent/internal/pkg/utils" - - "github.com/cloudwego/eino/adk" - einotool "github.com/cloudwego/eino/components/tool" - "github.com/cloudwego/eino/schema" - "github.com/google/uuid" + "cs-agent/internal/ai/runtime/internal/executor" ) type Service struct { - agentFactory *factory.AgentFactory - runnerFactory *factory.RunnerFactory + executor *executor.Service } func NewService() *Service { return &Service{ - agentFactory: factory.NewAgentFactory(), - runnerFactory: factory.NewRunnerFactory(), + executor: executor.NewService(), } } @@ -38,151 +21,7 @@ func (s *Service) Run(ctx context.Context, req Request) (*Summary, error) { } func (s *Service) ExecuteRun(ctx context.Context, req RunInput) (*RunResult, error) { - summary := &Summary{ - RunID: uuid.NewString(), - Status: "started", - ToolCodes: make([]string, 0), - InvokedToolCodes: make([]string, 0), - } - collector := callbacks.NewRuntimeTraceCollector() - collector.Data.RunID = summary.RunID - if req.AIAgent == nil || req.Conversation == nil || req.UserMessage == nil { - summary.Status = "error" - summary.ErrorMessage = "invalid runtime request" - collector.Data.Status = summary.Status - collector.Data.Error.Message = summary.ErrorMessage - collector.Data.Error.Stage = "prepare" - summary.TraceData = collector.Marshal() - return summary, fmt.Errorf("%s", summary.ErrorMessage) - } - if req.AIConfig == nil { - summary.Status = "error" - summary.ErrorMessage = "ai config is nil" - collector.Data.Status = summary.Status - collector.Data.Error.Message = summary.ErrorMessage - collector.Data.Error.Stage = "prepare" - summary.TraceData = collector.Marshal() - return summary, fmt.Errorf("%s", summary.ErrorMessage) - } - - history := adapter.BuildHistoryMessages(req.Conversation.ID, req.UserMessage.ID, 12) - summary.HistoryMessageCount = len(history.Messages) - collector.Data.Input.HistoryMessageCount = len(history.Messages) - collector.Data.Input.KnowledgeBaseIDs = utils.SplitInt64s(req.AIAgent.KnowledgeIDs) - collector.Data.Input.CurrentUserMessagePreview = preview(req.UserMessage.Content, 120) - - toolDefs, err := factory.NewToolFactory().BuildMCPTools(req.AIAgent) - if err != nil { - summary.Status = "error" - summary.ErrorMessage = err.Error() - collector.Data.Status = summary.Status - collector.Data.Error.Message = err.Error() - collector.Data.Error.Stage = "prepare" - summary.TraceData = collector.Marshal() - return summary, err - } - filteredToolDefs := filterToolDefinitionsBySkill(toolDefs, req.SelectedSkill) - toolDefsByModelName := make(map[string]string, len(filteredToolDefs)) - for _, item := range filteredToolDefs { - summary.ToolCodes = append(summary.ToolCodes, item.ToolCode) - toolDefsByModelName[item.ModelName] = item.ToolCode - } - if len(filteredToolDefs) > 0 { - summary.ToolCodes = appendIfMissing(summary.ToolCodes, toolx.BuiltinToolSearch.Code) - toolDefsByModelName[toolx.BuiltinToolSearch.Name] = toolx.BuiltinToolSearch.Code - } - if req.SelectedSkill != nil { - summary.ToolCodes = appendIfMissing(summary.ToolCodes, toolx.BuiltinSkill.Code) - toolDefsByModelName[toolx.BuiltinSkill.Name] = toolx.BuiltinSkill.Code - } - for modelName, toolCode := range toolSetStaticToolCodes(req.ToolSet) { - toolCode = strings.TrimSpace(toolCode) - modelName = strings.TrimSpace(modelName) - if toolCode == "" || modelName == "" { - continue - } - summary.ToolCodes = appendIfMissing(summary.ToolCodes, toolCode) - toolDefsByModelName[modelName] = toolCode - } - collector.Data.Input.ToolCodes = append(collector.Data.Input.ToolCodes, summary.ToolCodes...) - collector.SetTooling(staticToolCodeList(req.ToolSet), definitionToolCodes(filteredToolDefs), len(filteredToolDefs) > 0) - - collector.Data.Model.Provider = string(req.AIConfig.Provider) - collector.Data.Model.Name = req.AIConfig.ModelName - summary.SelectedSkillCode = "" - summary.SelectedSkillName = "" - summary.SkillRouteReason = strings.TrimSpace(req.SkillRouteReason) - summary.SkillRouteTrace = strings.TrimSpace(req.SkillRouteTrace) - if req.SelectedSkill != nil { - summary.SelectedSkillCode = strings.TrimSpace(req.SelectedSkill.Code) - summary.SelectedSkillName = strings.TrimSpace(req.SelectedSkill.Name) - summary.SkillAllowedToolCodes = parseJSONArrayList(req.SelectedSkill.ToolWhitelist) - collector.Data.Skill.Code = summary.SelectedSkillCode - collector.Data.Skill.Name = summary.SelectedSkillName - collector.Data.Skill.AllowedToolCodes = append([]string(nil), summary.SkillAllowedToolCodes...) - } - collector.Data.Skill.RouteReason = summary.SkillRouteReason - collector.Data.Skill.RouteTrace = summary.SkillRouteTrace - - agent, err := s.agentFactory.BuildCustomerServiceAgent(ctx, factory.BuildCustomerServiceAgentInput{ - AIAgent: req.AIAgent, - AIConfig: req.AIConfig, - SelectedSkill: req.SelectedSkill, - InstructionToolDefinitions: filteredToolDefs, - DynamicMCPToolDefinitions: filteredToolDefs, - StaticTools: toolSetStaticTools(req.ToolSet), - StaticToolCodes: toolSetStaticToolCodes(req.ToolSet), - Collector: collector, - }) - if err != nil { - summary.Status = "error" - summary.ErrorMessage = err.Error() - collector.Data.Status = summary.Status - collector.Data.Error.Message = err.Error() - collector.Data.Error.Stage = "prepare" - summary.TraceData = collector.Marshal() - return summary, err - } - - checkPointID := strings.TrimSpace(req.CheckPointID) - if checkPointID == "" { - checkPointID = "eino_cp_" + summary.RunID - } - summary.CheckPointID = checkPointID - runner := s.runnerFactory.Build(ctx, agent, false, true) - if runner == nil { - summary.Status = "error" - summary.ErrorMessage = "failed to build runner" - collector.Data.Status = summary.Status - collector.Data.Error.Message = summary.ErrorMessage - collector.Data.Error.Stage = "prepare" - summary.TraceData = collector.Marshal() - return summary, fmt.Errorf("%s", summary.ErrorMessage) - } - messages := make([]*schema.Message, 0, len(history.Messages)+3) - messages = append(messages, history.Messages...) - - retriever := retrievers.NewKnowledgeRetriever(req.AIAgent) - retrieveOptions := retrievers.DefaultKnowledgeRetrieveOptions() - retrieveOptions.QueryPreview = preview(req.UserMessage.Content, 120) - if retrieveResult, retrieveErr := retriever.RetrieveContextByOptions(ctx, retrieveOptions, strings.TrimSpace(req.UserMessage.Content)); retrieveErr == nil && retrieveResult != nil { - summary.RetrieverCount = len(retrieveResult.Hits) - collector.SetRetrieverSummary(retrieveResult.TraceSummary) - collector.Data.Retriever.Items = append(collector.Data.Retriever.Items, retrieveResult.TraceItems...) - if strings.TrimSpace(retrieveResult.ContextText) != "" { - messages = append(messages, schema.SystemMessage(retrieveResult.ContextText)) - } - } - - messages = append(messages, schema.UserMessage(strings.TrimSpace(req.UserMessage.Content))) - collector.Data.Interrupt.CheckPointID = checkPointID - consumeAgentEvents(runner.Run(ctx, messages, buildRunOptions(checkPointID)...), summary, collector, toolDefsByModelName) - 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 + return s.executor.ExecuteRun(ctx, executor.RunInput(req)) } func (s *Service) Resume(ctx context.Context, req ResumeRequest) (*Summary, error) { @@ -190,383 +29,5 @@ func (s *Service) Resume(ctx context.Context, req ResumeRequest) (*Summary, erro } func (s *Service) ExecuteResume(ctx context.Context, req ResumeInput) (*RunResult, error) { - summary := &Summary{ - RunID: uuid.NewString(), - Status: "started", - CheckPointID: strings.TrimSpace(req.CheckPointID), - ToolCodes: make([]string, 0), - InvokedToolCodes: make([]string, 0), - Interrupts: make([]InterruptContextSummary, 0), - } - collector := callbacks.NewRuntimeTraceCollector() - collector.Data.RunID = summary.RunID - collector.Data.Interrupt.CheckPointID = summary.CheckPointID - if req.AIAgent == nil { - summary.Status = "error" - summary.ErrorMessage = "ai agent is nil" - collector.Data.Status = summary.Status - collector.Data.Error.Message = summary.ErrorMessage - collector.Data.Error.Stage = "resume_prepare" - summary.TraceData = collector.Marshal() - return summary, fmt.Errorf("%s", summary.ErrorMessage) - } - if req.AIConfig == nil { - summary.Status = "error" - summary.ErrorMessage = "ai config is nil" - collector.Data.Status = summary.Status - collector.Data.Error.Message = summary.ErrorMessage - collector.Data.Error.Stage = "resume_prepare" - summary.TraceData = collector.Marshal() - return summary, fmt.Errorf("%s", summary.ErrorMessage) - } - if summary.CheckPointID == "" { - summary.Status = "error" - summary.ErrorMessage = "checkpoint id is required" - collector.Data.Status = summary.Status - collector.Data.Error.Message = summary.ErrorMessage - collector.Data.Error.Stage = "resume_prepare" - summary.TraceData = collector.Marshal() - return summary, fmt.Errorf("%s", summary.ErrorMessage) - } - toolDefs, err := factory.NewToolFactory().BuildMCPTools(req.AIAgent) - if err != nil { - summary.Status = "error" - summary.ErrorMessage = err.Error() - collector.Data.Status = summary.Status - collector.Data.Error.Message = err.Error() - collector.Data.Error.Stage = "resume_prepare" - summary.TraceData = collector.Marshal() - return summary, err - } - toolDefsByModelName := make(map[string]string, len(toolDefs)) - for _, item := range toolDefs { - summary.ToolCodes = append(summary.ToolCodes, item.ToolCode) - toolDefsByModelName[item.ModelName] = item.ToolCode - } - if len(toolDefs) > 0 { - summary.ToolCodes = appendIfMissing(summary.ToolCodes, toolx.BuiltinToolSearch.Code) - toolDefsByModelName[toolx.BuiltinToolSearch.Name] = toolx.BuiltinToolSearch.Code - } - for modelName, toolCode := range toolSetStaticToolCodes(req.ToolSet) { - toolCode = strings.TrimSpace(toolCode) - modelName = strings.TrimSpace(modelName) - if toolCode == "" || modelName == "" { - continue - } - summary.ToolCodes = appendIfMissing(summary.ToolCodes, toolCode) - toolDefsByModelName[modelName] = toolCode - } - collector.Data.Input.ToolCodes = append(collector.Data.Input.ToolCodes, summary.ToolCodes...) - collector.SetTooling(staticToolCodeList(req.ToolSet), definitionToolCodes(toolDefs), len(toolDefs) > 0) - collector.Data.Model.Provider = string(req.AIConfig.Provider) - collector.Data.Model.Name = req.AIConfig.ModelName - agent, err := s.agentFactory.BuildCustomerServiceAgent(ctx, factory.BuildCustomerServiceAgentInput{ - AIAgent: req.AIAgent, - AIConfig: req.AIConfig, - InstructionToolDefinitions: toolDefs, - DynamicMCPToolDefinitions: toolDefs, - StaticTools: toolSetStaticTools(req.ToolSet), - StaticToolCodes: toolSetStaticToolCodes(req.ToolSet), - Collector: collector, - }) - if err != nil { - summary.Status = "error" - summary.ErrorMessage = err.Error() - collector.Data.Status = summary.Status - collector.Data.Error.Message = err.Error() - collector.Data.Error.Stage = "resume_prepare" - summary.TraceData = collector.Marshal() - return summary, err - } - runner := s.runnerFactory.Build(ctx, agent, false, true) - if runner == nil { - summary.Status = "error" - summary.ErrorMessage = "failed to build runner" - collector.Data.Status = summary.Status - collector.Data.Error.Message = summary.ErrorMessage - collector.Data.Error.Stage = "resume_prepare" - summary.TraceData = collector.Marshal() - return summary, fmt.Errorf("%s", summary.ErrorMessage) - } - var iter *adk.AsyncIterator[*adk.AgentEvent] - if len(req.ResumeData) > 0 { - iter, err = runner.ResumeWithParams(ctx, summary.CheckPointID, &adk.ResumeParams{Targets: req.ResumeData}) - } else { - iter, err = runner.Resume(ctx, summary.CheckPointID) - } - if err != nil { - summary.Status = "error" - summary.ErrorMessage = err.Error() - collector.Data.Status = summary.Status - collector.Data.Error.Message = err.Error() - collector.Data.Error.Stage = "resume_prepare" - summary.TraceData = collector.Marshal() - return summary, err - } - consumeAgentEvents(iter, summary, collector, toolDefsByModelName) - 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 -} - -func filterToolDefinitionsBySkill(definitions []adapter.MCPToolDefinition, skill *models.SkillDefinition) []adapter.MCPToolDefinition { - if len(definitions) == 0 || skill == nil { - return definitions - } - allowed := parseJSONArraySet(skill.ToolWhitelist) - if len(allowed) == 0 { - return definitions - } - ret := make([]adapter.MCPToolDefinition, 0, len(definitions)) - for _, item := range definitions { - if _, ok := allowed[strings.TrimSpace(item.ToolCode)]; ok { - ret = append(ret, item) - } - } - return ret -} - -func parseJSONArraySet(raw string) map[string]struct{} { - raw = strings.TrimSpace(raw) - if raw == "" { - return nil - } - items := parseJSONArrayList(raw) - if len(items) == 0 { - return nil - } - ret := make(map[string]struct{}, len(items)) - for _, item := range items { - ret[item] = struct{}{} - } - return ret -} - -func parseJSONArrayList(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) - } - return ret -} - -func buildRunOptions(checkPointID string) []adk.AgentRunOption { - if strings.TrimSpace(checkPointID) == "" { - return nil - } - return []adk.AgentRunOption{adk.WithCheckPointID(checkPointID)} -} - -func consumeAgentEvents(iter *adk.AsyncIterator[*adk.AgentEvent], summary *Summary, collector *callbacks.RuntimeTraceCollector, toolDefsByModelName map[string]string) { - if iter == nil || summary == nil || collector == nil { - return - } - for { - event, ok := iter.Next() - if !ok { - break - } - if event == nil { - continue - } - if event.Err != nil { - summary.Status = "error" - summary.ErrorMessage = event.Err.Error() - collector.Data.Error.Message = event.Err.Error() - collector.Data.Error.Stage = "model" - continue - } - if event.Action != nil && event.Action.Interrupted != nil { - summary.Status = "interrupted" - summary.Interrupted = true - summary.Interrupts = summarizeInterrupts(event.Action.Interrupted.InterruptContexts) - collector.Data.Interrupt.Items = convertInterruptTraceItems(summary.Interrupts) - continue - } - if event.Output == nil || event.Output.MessageOutput == nil { - continue - } - message, getErr := event.Output.MessageOutput.GetMessage() - if getErr != nil || message == nil { - continue - } - switch event.Output.MessageOutput.Role { - case schema.Assistant: - summary.ReplyText = strings.TrimSpace(message.Content) - case schema.Tool: - summary.ToolCallCount++ - if toolDefsByModelName != nil { - toolCode := strings.TrimSpace(toolDefsByModelName[message.ToolName]) - if toolCode != "" { - summary.InvokedToolCodes = appendIfMissing(summary.InvokedToolCodes, toolCode) - } - } - } - } - if summary.Status == "started" { - if strings.TrimSpace(summary.ReplyText) == "" { - summary.Status = "fallback" - } else { - summary.Status = "completed" - } - } -} - -func convertInterruptTraceItems(items []InterruptContextSummary) []callbacks.InterruptTraceContext { - if len(items) == 0 { - return nil - } - ret := make([]callbacks.InterruptTraceContext, 0, len(items)) - for _, item := range items { - ret = append(ret, callbacks.InterruptTraceContext{ - Type: item.Type, - ID: item.ID, - InfoPreview: item.InfoPreview, - }) - } - return ret -} - -func previewInterruptInfo(info any) string { - if info == nil { - return "" - } - switch v := info.(type) { - case string: - return preview(v, 200) - default: - data, err := json.Marshal(v) - if err != nil { - return "" - } - return preview(string(data), 200) - } -} - -func summarizeInterrupts(items []*adk.InterruptCtx) []InterruptContextSummary { - if len(items) == 0 { - return nil - } - ret := make([]InterruptContextSummary, 0, len(items)) - for _, item := range items { - if item == nil { - continue - } - ret = append(ret, InterruptContextSummary{ - Type: extractInterruptType(item.Info), - ID: strings.TrimSpace(item.ID), - InfoPreview: previewInterruptInfo(item.Info), - }) - } - return ret -} - -func extractInterruptType(info any) string { - if info == nil { - return "" - } - switch v := info.(type) { - case map[string]any: - return strings.TrimSpace(getStringFromAnyMap(v, "type")) - default: - return "" - } -} - -func getStringFromAnyMap(data map[string]any, key string) string { - value, ok := data[key] - if !ok || value == nil { - return "" - } - switch v := value.(type) { - case string: - return v - default: - return fmt.Sprintf("%v", v) - } -} - -func appendIfMissing(items []string, value string) []string { - value = strings.TrimSpace(value) - if value == "" { - return items - } - for _, item := range items { - if strings.TrimSpace(item) == value { - return items - } - } - return append(items, value) -} - -func preview(value string, limit int) string { - if limit <= 0 { - return "" - } - value = strings.TrimSpace(value) - runes := []rune(value) - if len(runes) <= limit { - return value - } - return string(runes[:limit]) + "..." -} - -func toolSetStaticTools(toolSet *registry.ToolSet) []einotool.BaseTool { - if toolSet == nil { - return nil - } - return toolSet.StaticTools -} - -func toolSetStaticToolCodes(toolSet *registry.ToolSet) map[string]string { - if toolSet == nil { - return nil - } - return toolSet.StaticToolCodes -} - -func definitionToolCodes(definitions []adapter.MCPToolDefinition) []string { - if len(definitions) == 0 { - return nil - } - ret := make([]string, 0, len(definitions)) - for _, item := range definitions { - toolCode := strings.TrimSpace(item.ToolCode) - if toolCode == "" { - continue - } - ret = append(ret, toolCode) - } - return ret -} - -func staticToolCodeList(toolSet *registry.ToolSet) []string { - toolCodes := toolSetStaticToolCodes(toolSet) - if len(toolCodes) == 0 { - return nil - } - ret := make([]string, 0, len(toolCodes)) - for _, toolCode := range toolCodes { - toolCode = strings.TrimSpace(toolCode) - if toolCode == "" { - continue - } - ret = append(ret, toolCode) - } - return ret + return s.executor.ExecuteResume(ctx, executor.ResumeInput(req)) } diff --git a/internal/ai/runtime/internal/engine/types.go b/internal/ai/runtime/internal/engine/types.go index 2945bec..e239405 100644 --- a/internal/ai/runtime/internal/engine/types.go +++ b/internal/ai/runtime/internal/engine/types.go @@ -1,60 +1,12 @@ package engine -import ( - "cs-agent/internal/ai/runtime/registry" - "cs-agent/internal/models" -) +import "cs-agent/internal/ai/runtime/internal/executor" -type RunInput struct { - Conversation *models.Conversation - UserMessage *models.Message - AIAgent *models.AIAgent - AIConfig *models.AIConfig - SelectedSkill *models.SkillDefinition - SkillRouteReason string - SkillRouteTrace string - CheckPointID string - ToolSet *registry.ToolSet -} - -type ResumeInput struct { - Conversation *models.Conversation - AIAgent *models.AIAgent - AIConfig *models.AIConfig - CheckPointID string - ResumeData map[string]any - ToolSet *registry.ToolSet -} - -type InterruptContextSummary struct { - Type string `json:"type,omitempty"` - ID string `json:"id"` - InfoPreview string `json:"infoPreview,omitempty"` -} - -type RunResult struct { - RunID string - Status string - ReplyText string - SelectedSkillCode string - SelectedSkillName string - SkillRouteReason string - SkillRouteTrace string - SkillAllowedToolCodes []string - ModelName string - PromptTokens int - CompletionTokens int - HistoryMessageCount int - RetrieverCount int - ToolCallCount int - ToolCodes []string - InvokedToolCodes []string - CheckPointID string - Interrupted bool - Interrupts []InterruptContextSummary - TraceData string - ErrorMessage string -} +// TODO 这个地方为什么要定义类型别名,不能直接用吗? +type RunInput = executor.RunInput +type ResumeInput = executor.ResumeInput +type InterruptContextSummary = executor.InterruptContextSummary +type RunResult = executor.RunResult type Request = RunInput type ResumeRequest = ResumeInput diff --git a/internal/ai/runtime/internal/executor/service.go b/internal/ai/runtime/internal/executor/service.go new file mode 100644 index 0000000..20062bd --- /dev/null +++ b/internal/ai/runtime/internal/executor/service.go @@ -0,0 +1,541 @@ +package executor + +import ( + "context" + "encoding/json" + "fmt" + "strings" + + "cs-agent/internal/ai/runtime/internal/impl/adapter" + "cs-agent/internal/ai/runtime/internal/impl/callbacks" + "cs-agent/internal/ai/runtime/internal/impl/factory" + "cs-agent/internal/ai/runtime/internal/impl/retrievers" + "cs-agent/internal/ai/runtime/registry" + "cs-agent/internal/models" + "cs-agent/internal/pkg/toolx" + "cs-agent/internal/pkg/utils" + + "github.com/cloudwego/eino/adk" + einotool "github.com/cloudwego/eino/components/tool" + "github.com/cloudwego/eino/schema" + "github.com/google/uuid" +) + +type Service struct { + agentFactory *factory.AgentFactory + runnerFactory *factory.RunnerFactory +} + +func NewService() *Service { + return &Service{ + agentFactory: factory.NewAgentFactory(), + runnerFactory: factory.NewRunnerFactory(), + } +} + +func (s *Service) ExecuteRun(ctx context.Context, req RunInput) (*RunResult, error) { + summary := &RunResult{ + RunID: uuid.NewString(), + Status: "started", + ToolCodes: make([]string, 0), + InvokedToolCodes: make([]string, 0), + } + collector := callbacks.NewRuntimeTraceCollector() + collector.Data.RunID = summary.RunID + if req.AIAgent == nil || req.Conversation == nil || req.UserMessage == nil { + summary.Status = "error" + summary.ErrorMessage = "invalid runtime request" + collector.Data.Status = summary.Status + collector.Data.Error.Message = summary.ErrorMessage + collector.Data.Error.Stage = "prepare" + summary.TraceData = collector.Marshal() + return summary, fmt.Errorf("%s", summary.ErrorMessage) + } + if req.AIConfig == nil { + summary.Status = "error" + summary.ErrorMessage = "ai config is nil" + collector.Data.Status = summary.Status + collector.Data.Error.Message = summary.ErrorMessage + collector.Data.Error.Stage = "prepare" + summary.TraceData = collector.Marshal() + return summary, fmt.Errorf("%s", summary.ErrorMessage) + } + + history := adapter.BuildHistoryMessages(req.Conversation.ID, req.UserMessage.ID, 12) + summary.HistoryMessageCount = len(history.Messages) + collector.Data.Input.HistoryMessageCount = len(history.Messages) + collector.Data.Input.KnowledgeBaseIDs = utils.SplitInt64s(req.AIAgent.KnowledgeIDs) + collector.Data.Input.CurrentUserMessagePreview = preview(req.UserMessage.Content, 120) + + toolDefs, err := factory.NewToolFactory().BuildMCPTools(req.AIAgent) + if err != nil { + summary.Status = "error" + summary.ErrorMessage = err.Error() + collector.Data.Status = summary.Status + collector.Data.Error.Message = err.Error() + collector.Data.Error.Stage = "prepare" + summary.TraceData = collector.Marshal() + return summary, err + } + filteredToolDefs := filterToolDefinitionsBySkill(toolDefs, req.SelectedSkill) + toolDefsByModelName := make(map[string]string, len(filteredToolDefs)) + for _, item := range filteredToolDefs { + summary.ToolCodes = append(summary.ToolCodes, item.ToolCode) + toolDefsByModelName[item.ModelName] = item.ToolCode + } + if len(filteredToolDefs) > 0 { + summary.ToolCodes = appendIfMissing(summary.ToolCodes, toolx.BuiltinToolSearch.Code) + toolDefsByModelName[toolx.BuiltinToolSearch.Name] = toolx.BuiltinToolSearch.Code + } + if req.SelectedSkill != nil { + summary.ToolCodes = appendIfMissing(summary.ToolCodes, toolx.BuiltinSkill.Code) + toolDefsByModelName[toolx.BuiltinSkill.Name] = toolx.BuiltinSkill.Code + } + for modelName, toolCode := range toolSetStaticToolCodes(req.ToolSet) { + toolCode = strings.TrimSpace(toolCode) + modelName = strings.TrimSpace(modelName) + if toolCode == "" || modelName == "" { + continue + } + summary.ToolCodes = appendIfMissing(summary.ToolCodes, toolCode) + toolDefsByModelName[modelName] = toolCode + } + collector.Data.Input.ToolCodes = append(collector.Data.Input.ToolCodes, summary.ToolCodes...) + collector.SetTooling(staticToolCodeList(req.ToolSet), definitionToolCodes(filteredToolDefs), len(filteredToolDefs) > 0) + + collector.Data.Model.Provider = string(req.AIConfig.Provider) + collector.Data.Model.Name = req.AIConfig.ModelName + summary.SelectedSkillCode = "" + summary.SelectedSkillName = "" + summary.SkillRouteReason = strings.TrimSpace(req.SkillRouteReason) + summary.SkillRouteTrace = strings.TrimSpace(req.SkillRouteTrace) + if req.SelectedSkill != nil { + summary.SelectedSkillCode = strings.TrimSpace(req.SelectedSkill.Code) + summary.SelectedSkillName = strings.TrimSpace(req.SelectedSkill.Name) + summary.SkillAllowedToolCodes = parseJSONArrayList(req.SelectedSkill.ToolWhitelist) + collector.Data.Skill.Code = summary.SelectedSkillCode + collector.Data.Skill.Name = summary.SelectedSkillName + collector.Data.Skill.AllowedToolCodes = append([]string(nil), summary.SkillAllowedToolCodes...) + } + collector.Data.Skill.RouteReason = summary.SkillRouteReason + collector.Data.Skill.RouteTrace = summary.SkillRouteTrace + + agent, err := s.agentFactory.BuildCustomerServiceAgent(ctx, factory.BuildCustomerServiceAgentInput{ + AIAgent: req.AIAgent, + AIConfig: req.AIConfig, + SelectedSkill: req.SelectedSkill, + InstructionToolDefinitions: filteredToolDefs, + DynamicMCPToolDefinitions: filteredToolDefs, + StaticTools: toolSetStaticTools(req.ToolSet), + StaticToolCodes: toolSetStaticToolCodes(req.ToolSet), + Collector: collector, + }) + if err != nil { + summary.Status = "error" + summary.ErrorMessage = err.Error() + collector.Data.Status = summary.Status + collector.Data.Error.Message = err.Error() + collector.Data.Error.Stage = "prepare" + summary.TraceData = collector.Marshal() + return summary, err + } + + checkPointID := strings.TrimSpace(req.CheckPointID) + if checkPointID == "" { + checkPointID = "eino_cp_" + summary.RunID + } + summary.CheckPointID = checkPointID + runner := s.runnerFactory.Build(ctx, agent, false, true) + if runner == nil { + summary.Status = "error" + summary.ErrorMessage = "failed to build runner" + collector.Data.Status = summary.Status + collector.Data.Error.Message = summary.ErrorMessage + collector.Data.Error.Stage = "prepare" + summary.TraceData = collector.Marshal() + return summary, fmt.Errorf("%s", summary.ErrorMessage) + } + messages := make([]*schema.Message, 0, len(history.Messages)+3) + messages = append(messages, history.Messages...) + + retriever := retrievers.NewKnowledgeRetriever(req.AIAgent) + retrieveOptions := retrievers.DefaultKnowledgeRetrieveOptions() + retrieveOptions.QueryPreview = preview(req.UserMessage.Content, 120) + if retrieveResult, retrieveErr := retriever.RetrieveContextByOptions(ctx, retrieveOptions, strings.TrimSpace(req.UserMessage.Content)); retrieveErr == nil && retrieveResult != nil { + summary.RetrieverCount = len(retrieveResult.Hits) + collector.SetRetrieverSummary(retrieveResult.TraceSummary) + collector.Data.Retriever.Items = append(collector.Data.Retriever.Items, retrieveResult.TraceItems...) + if strings.TrimSpace(retrieveResult.ContextText) != "" { + messages = append(messages, schema.SystemMessage(retrieveResult.ContextText)) + } + } + + messages = append(messages, schema.UserMessage(strings.TrimSpace(req.UserMessage.Content))) + collector.Data.Interrupt.CheckPointID = checkPointID + consumeAgentEvents(runner.Run(ctx, messages, buildRunOptions(checkPointID)...), summary, collector, toolDefsByModelName) + 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 +} + +func (s *Service) ExecuteResume(ctx context.Context, req ResumeInput) (*RunResult, error) { + summary := &RunResult{ + RunID: uuid.NewString(), + Status: "started", + CheckPointID: strings.TrimSpace(req.CheckPointID), + ToolCodes: make([]string, 0), + InvokedToolCodes: make([]string, 0), + Interrupts: make([]InterruptContextSummary, 0), + } + collector := callbacks.NewRuntimeTraceCollector() + collector.Data.RunID = summary.RunID + collector.Data.Interrupt.CheckPointID = summary.CheckPointID + if req.AIAgent == nil { + summary.Status = "error" + summary.ErrorMessage = "ai agent is nil" + collector.Data.Status = summary.Status + collector.Data.Error.Message = summary.ErrorMessage + collector.Data.Error.Stage = "resume_prepare" + summary.TraceData = collector.Marshal() + return summary, fmt.Errorf("%s", summary.ErrorMessage) + } + if req.AIConfig == nil { + summary.Status = "error" + summary.ErrorMessage = "ai config is nil" + collector.Data.Status = summary.Status + collector.Data.Error.Message = summary.ErrorMessage + collector.Data.Error.Stage = "resume_prepare" + summary.TraceData = collector.Marshal() + return summary, fmt.Errorf("%s", summary.ErrorMessage) + } + if summary.CheckPointID == "" { + summary.Status = "error" + summary.ErrorMessage = "checkpoint id is required" + collector.Data.Status = summary.Status + collector.Data.Error.Message = summary.ErrorMessage + collector.Data.Error.Stage = "resume_prepare" + summary.TraceData = collector.Marshal() + return summary, fmt.Errorf("%s", summary.ErrorMessage) + } + toolDefs, err := factory.NewToolFactory().BuildMCPTools(req.AIAgent) + if err != nil { + summary.Status = "error" + summary.ErrorMessage = err.Error() + collector.Data.Status = summary.Status + collector.Data.Error.Message = err.Error() + collector.Data.Error.Stage = "resume_prepare" + summary.TraceData = collector.Marshal() + return summary, err + } + toolDefsByModelName := make(map[string]string, len(toolDefs)) + for _, item := range toolDefs { + summary.ToolCodes = append(summary.ToolCodes, item.ToolCode) + toolDefsByModelName[item.ModelName] = item.ToolCode + } + if len(toolDefs) > 0 { + summary.ToolCodes = appendIfMissing(summary.ToolCodes, toolx.BuiltinToolSearch.Code) + toolDefsByModelName[toolx.BuiltinToolSearch.Name] = toolx.BuiltinToolSearch.Code + } + for modelName, toolCode := range toolSetStaticToolCodes(req.ToolSet) { + toolCode = strings.TrimSpace(toolCode) + modelName = strings.TrimSpace(modelName) + if toolCode == "" || modelName == "" { + continue + } + summary.ToolCodes = appendIfMissing(summary.ToolCodes, toolCode) + toolDefsByModelName[modelName] = toolCode + } + collector.Data.Input.ToolCodes = append(collector.Data.Input.ToolCodes, summary.ToolCodes...) + collector.SetTooling(staticToolCodeList(req.ToolSet), definitionToolCodes(toolDefs), len(toolDefs) > 0) + collector.Data.Model.Provider = string(req.AIConfig.Provider) + collector.Data.Model.Name = req.AIConfig.ModelName + + agent, err := s.agentFactory.BuildCustomerServiceAgent(ctx, factory.BuildCustomerServiceAgentInput{ + AIAgent: req.AIAgent, + AIConfig: req.AIConfig, + InstructionToolDefinitions: toolDefs, + DynamicMCPToolDefinitions: toolDefs, + StaticTools: toolSetStaticTools(req.ToolSet), + StaticToolCodes: toolSetStaticToolCodes(req.ToolSet), + Collector: collector, + }) + if err != nil { + summary.Status = "error" + summary.ErrorMessage = err.Error() + collector.Data.Status = summary.Status + collector.Data.Error.Message = err.Error() + collector.Data.Error.Stage = "resume_prepare" + summary.TraceData = collector.Marshal() + return summary, err + } + runner := s.runnerFactory.Build(ctx, agent, false, true) + if runner == nil { + summary.Status = "error" + summary.ErrorMessage = "failed to build runner" + collector.Data.Status = summary.Status + collector.Data.Error.Message = summary.ErrorMessage + collector.Data.Error.Stage = "resume_prepare" + summary.TraceData = collector.Marshal() + return summary, fmt.Errorf("%s", summary.ErrorMessage) + } + resumeData := buildResumeDataMessage(req.ResumeData) + iter, err := runner.Resume(ctx, summary.CheckPointID, buildResumeOptions(summary.CheckPointID, resumeData)...) + if err != nil { + summary.Status = "error" + summary.ErrorMessage = err.Error() + collector.Data.Status = summary.Status + collector.Data.Error.Message = err.Error() + collector.Data.Error.Stage = "resume_execute" + summary.TraceData = collector.Marshal() + return summary, err + } + consumeAgentEvents(iter, summary, collector, toolDefsByModelName) + 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 +} + +func buildResumeDataMessage(resumeData map[string]any) *schema.Message { + if len(resumeData) == 0 { + return nil + } + data, err := json.Marshal(resumeData) + if err != nil { + return schema.UserMessage(fmt.Sprint(resumeData)) + } + return schema.UserMessage(string(data)) +} + +func buildRunOptions(checkPointID string) []adk.AgentRunOption { + options := make([]adk.AgentRunOption, 0, 1) + if strings.TrimSpace(checkPointID) != "" { + options = append(options, adk.WithCheckPointID(checkPointID)) + } + return options +} + +func buildResumeOptions(checkPointID string, resumeData *schema.Message) []adk.AgentRunOption { + options := make([]adk.AgentRunOption, 0, 1) + if strings.TrimSpace(checkPointID) != "" { + options = append(options, adk.WithCheckPointID(checkPointID)) + } + _ = resumeData + return options +} + +func consumeAgentEvents(events *adk.AsyncIterator[*adk.AgentEvent], summary *RunResult, collector *callbacks.RuntimeTraceCollector, toolDefsByModelName map[string]string) { + if summary == nil { + return + } + if collector == nil { + collector = callbacks.NewRuntimeTraceCollector() + } + for { + event, ok := events.Next() + if !ok { + break + } + if event == nil { + continue + } + if event.Action != nil && event.Action.Interrupted != nil { + summary.Status = "interrupted" + summary.Interrupted = true + summary.Interrupts = buildInterruptSummaries(event) + } + if event.Err != nil { + errMsg := strings.TrimSpace(event.Err.Error()) + if errMsg != "" { + summary.Status = "error" + summary.ErrorMessage = errMsg + } + } + if event.Output == nil || event.Output.MessageOutput == nil { + continue + } + messageOutput := event.Output.MessageOutput + switch messageOutput.Role { + case schema.Assistant: + replyText := strings.TrimSpace(messageOutput.Message.Content) + if replyText != "" { + summary.ReplyText = replyText + } + case schema.Tool: + toolName := strings.TrimSpace(messageOutput.ToolName) + if toolName == "" { + continue + } + toolCode := toolName + if mappedCode, ok := toolDefsByModelName[toolName]; ok && strings.TrimSpace(mappedCode) != "" { + toolCode = strings.TrimSpace(mappedCode) + } + summary.InvokedToolCodes = appendIfMissing(summary.InvokedToolCodes, toolCode) + } + } + if summary.Status == "started" { + switch { + case strings.TrimSpace(summary.ErrorMessage) != "": + summary.Status = "error" + case summary.Interrupted: + summary.Status = "interrupted" + case strings.TrimSpace(summary.ReplyText) != "": + summary.Status = "completed" + default: + summary.Status = "fallback" + } + } + summary.ToolCallCount = len(summary.InvokedToolCodes) +} + +func buildInterruptSummaries(event *adk.AgentEvent) []InterruptContextSummary { + if event == nil || event.Action == nil || event.Action.Interrupted == nil { + return nil + } + interrupts := event.Action.Interrupted.InterruptContexts + result := make([]InterruptContextSummary, 0, len(interrupts)) + for _, item := range interrupts { + if item == nil { + continue + } + result = append(result, InterruptContextSummary{ + ID: strings.TrimSpace(item.ID), + InfoPreview: previewInterruptInfo(item.Info), + }) + } + return result +} + +func previewInterruptInfo(info any) string { + if info == nil { + return "" + } + data, err := json.Marshal(info) + if err != nil { + return fmt.Sprint(info) + } + return string(data) +} + +func appendIfMissing(items []string, item string) []string { + item = strings.TrimSpace(item) + if item == "" { + return items + } + for _, existing := range items { + if strings.TrimSpace(existing) == item { + return items + } + } + return append(items, item) +} + +func staticToolCodeList(toolSet *registry.ToolSet) []string { + if toolSet == nil || len(toolSet.StaticToolCodes) == 0 { + return nil + } + ret := make([]string, 0, len(toolSet.StaticToolCodes)) + for _, code := range toolSet.StaticToolCodes { + code = strings.TrimSpace(code) + if code == "" { + continue + } + ret = append(ret, code) + } + return ret +} + +func toolSetStaticTools(toolSet *registry.ToolSet) []einotool.BaseTool { + if toolSet == nil { + return nil + } + return append([]einotool.BaseTool(nil), toolSet.StaticTools...) +} + +func toolSetStaticToolCodes(toolSet *registry.ToolSet) map[string]string { + if toolSet == nil || len(toolSet.StaticToolCodes) == 0 { + return nil + } + ret := make(map[string]string, len(toolSet.StaticToolCodes)) + for name, code := range toolSet.StaticToolCodes { + ret[strings.TrimSpace(name)] = strings.TrimSpace(code) + } + return ret +} + +func definitionToolCodes(defs []adapter.MCPToolDefinition) []string { + ret := make([]string, 0, len(defs)) + for _, item := range defs { + code := strings.TrimSpace(item.ToolCode) + if code == "" { + continue + } + ret = append(ret, code) + } + return ret +} + +func filterToolDefinitionsBySkill(defs []adapter.MCPToolDefinition, skill *models.SkillDefinition) []adapter.MCPToolDefinition { + if skill == nil || strings.TrimSpace(skill.ToolWhitelist) == "" { + return defs + } + var allowed []string + if err := json.Unmarshal([]byte(skill.ToolWhitelist), &allowed); err != nil { + return defs + } + allowedSet := make(map[string]struct{}, len(allowed)) + for _, item := range allowed { + item = toolx.NormalizeToolCodeAlias(item) + if strings.TrimSpace(item) == "" { + continue + } + allowedSet[strings.TrimSpace(item)] = struct{}{} + } + if len(allowedSet) == 0 { + return defs + } + ret := make([]adapter.MCPToolDefinition, 0, len(defs)) + for _, item := range defs { + if _, ok := allowedSet[strings.TrimSpace(item.ToolCode)]; ok { + ret = append(ret, item) + } + } + return ret +} + +func parseJSONArrayList(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) + } + return ret +} + +func preview(text string, limit int) string { + text = strings.TrimSpace(text) + if text == "" || limit <= 0 { + return "" + } + runes := []rune(text) + if len(runes) <= limit { + return string(runes) + } + return string(runes[:limit]) + "..." +} diff --git a/internal/ai/runtime/internal/executor/types.go b/internal/ai/runtime/internal/executor/types.go new file mode 100644 index 0000000..9b35a38 --- /dev/null +++ b/internal/ai/runtime/internal/executor/types.go @@ -0,0 +1,57 @@ +package executor + +import ( + "cs-agent/internal/ai/runtime/registry" + "cs-agent/internal/models" +) + +type RunInput struct { + Conversation *models.Conversation + UserMessage *models.Message + AIAgent *models.AIAgent + AIConfig *models.AIConfig + SelectedSkill *models.SkillDefinition + SkillRouteReason string + SkillRouteTrace string + CheckPointID string + ToolSet *registry.ToolSet +} + +type ResumeInput struct { + Conversation *models.Conversation + AIAgent *models.AIAgent + AIConfig *models.AIConfig + CheckPointID string + ResumeData map[string]any + ToolSet *registry.ToolSet +} + +type InterruptContextSummary struct { + Type string `json:"type,omitempty"` + ID string `json:"id"` + InfoPreview string `json:"infoPreview,omitempty"` +} + +type RunResult struct { + RunID string + Status string + ReplyText string + SelectedSkillCode string + SelectedSkillName string + SkillRouteReason string + SkillRouteTrace string + SkillAllowedToolCodes []string + ModelName string + PromptTokens int + CompletionTokens int + HistoryMessageCount int + RetrieverCount int + ToolCallCount int + ToolCodes []string + InvokedToolCodes []string + CheckPointID string + Interrupted bool + Interrupts []InterruptContextSummary + TraceData string + ErrorMessage string +}