diff --git a/internal/ai/runtime/internal/impl/factory/tool_factory.go b/internal/ai/runtime/internal/impl/factory/tool_factory.go index ad5d76d..82c5fed 100644 --- a/internal/ai/runtime/internal/impl/factory/tool_factory.go +++ b/internal/ai/runtime/internal/impl/factory/tool_factory.go @@ -2,13 +2,12 @@ package factory import ( "context" - "encoding/json" "strings" "cs-agent/internal/ai/mcps" impladapter "cs-agent/internal/ai/runtime/internal/impl/adapter" "cs-agent/internal/models" - "cs-agent/internal/pkg/dto/request" + "cs-agent/internal/pkg/toolx" einotool "github.com/cloudwego/eino/components/tool" ) @@ -23,13 +22,16 @@ func (f *ToolFactory) BuildMCPTools(aiAgent *models.AIAgent) ([]impladapter.MCPT if aiAgent == nil || strings.TrimSpace(aiAgent.AllowedMCPTools) == "" { return nil, nil } - var raw []request.AIAgentMCPToolRequest - if err := json.Unmarshal([]byte(aiAgent.AllowedMCPTools), &raw); err != nil { + raw, err := toolx.ParseAgentMCPToolsJSON(aiAgent.AllowedMCPTools) + if err != nil { return nil, err } ret := make([]impladapter.MCPToolDefinition, 0, len(raw)) for _, item := range raw { - toolCode := strings.TrimSpace(item.ServerCode) + "/" + strings.TrimSpace(item.ToolName) + toolCode := strings.TrimSpace(item.ToolCode) + if toolCode == "" { + toolCode = toolx.BuildMCPToolCode(item.ServerCode, item.ToolName) + } definition := impladapter.MCPToolDefinition{ ToolCode: toolCode, ServerCode: strings.TrimSpace(item.ServerCode), @@ -80,7 +82,7 @@ func (f *ToolFactory) loadToolMetadata(ctx context.Context, definitions []implad } for i := range toolInfos { toolInfo := toolInfos[i] - toolCode := strings.TrimSpace(serverCode) + "/" + strings.TrimSpace(toolInfo.Name) + toolCode := toolx.BuildMCPToolCode(serverCode, toolInfo.Name) toolInfoCopy := toolInfo toolsByCode[toolCode] = &toolInfoCopy } diff --git a/internal/controllers/console/ai_agent_controller.go b/internal/controllers/console/ai_agent_controller.go index 35b3b22..1be2b9d 100644 --- a/internal/controllers/console/ai_agent_controller.go +++ b/internal/controllers/console/ai_agent_controller.go @@ -8,6 +8,7 @@ import ( "cs-agent/internal/pkg/dto/request" "cs-agent/internal/pkg/dto/response" "cs-agent/internal/pkg/enums" + "cs-agent/internal/pkg/toolx" "cs-agent/internal/pkg/utils" "cs-agent/internal/services" @@ -117,15 +118,15 @@ func buildAIAgentResponse(item *models.AIAgent) response.AIAgentResponse { StatusName: enums.GetStatusLabel(item.Status), AIConfigID: item.AIConfigID, ServiceMode: item.ServiceMode, - ServiceModeName: enums.GetIMConversationServiceModeLabel(enums.IMConversationServiceMode(item.ServiceMode)), + ServiceModeName: enums.GetIMConversationServiceModeLabel(item.ServiceMode), SystemPrompt: item.SystemPrompt, WelcomeMessage: item.WelcomeMessage, ReplyTimeoutSeconds: item.ReplyTimeoutSeconds, HandoffMode: item.HandoffMode, - HandoffModeName: enums.GetAIAgentHandoffModeLabel(enums.AIAgentHandoffMode(item.HandoffMode)), + HandoffModeName: enums.GetAIAgentHandoffModeLabel(item.HandoffMode), MaxAIReplyRounds: item.MaxAIReplyRounds, FallbackMode: item.FallbackMode, - FallbackModeName: enums.GetAIAgentFallbackModeLabel(enums.AIAgentFallbackMode(item.FallbackMode)), + FallbackModeName: enums.GetAIAgentFallbackModeLabel(item.FallbackMode), FallbackMessage: item.FallbackMessage, KnowledgeIDs: utils.SplitInt64s(item.KnowledgeIDs), SkillIDs: utils.SplitInt64s(item.SkillIDs), @@ -169,7 +170,12 @@ func buildAIAgentResponse(item *models.AIAgent) response.AIAgentResponse { var directTools []request.AIAgentMCPToolRequest if err := json.Unmarshal([]byte(raw), &directTools); err == nil { for _, tool := range directTools { + toolCode := strings.TrimSpace(tool.ToolCode) + if toolCode == "" { + toolCode = toolx.BuildMCPToolCode(tool.ServerCode, tool.ToolName) + } ret.DirectTools = append(ret.DirectTools, response.AIAgentMCPToolResponse{ + ToolCode: toolCode, ServerCode: strings.TrimSpace(tool.ServerCode), ToolName: strings.TrimSpace(tool.ToolName), Title: strings.TrimSpace(tool.Title), diff --git a/internal/pkg/dto/request/ai_request.go b/internal/pkg/dto/request/ai_request.go index c6ffe4c..8538179 100644 --- a/internal/pkg/dto/request/ai_request.go +++ b/internal/pkg/dto/request/ai_request.go @@ -3,6 +3,7 @@ package request import "cs-agent/internal/pkg/enums" type AIAgentMCPToolRequest struct { + ToolCode string `json:"toolCode"` ServerCode string `json:"serverCode"` ToolName string `json:"toolName"` Title string `json:"title"` diff --git a/internal/pkg/dto/response/ai_response.go b/internal/pkg/dto/response/ai_response.go index 2e97385..4d82516 100644 --- a/internal/pkg/dto/response/ai_response.go +++ b/internal/pkg/dto/response/ai_response.go @@ -17,6 +17,7 @@ type AIAgentSkillResponse struct { } type AIAgentMCPToolResponse struct { + ToolCode string `json:"toolCode"` ServerCode string `json:"serverCode"` ToolName string `json:"toolName"` Title string `json:"title"` diff --git a/internal/pkg/toolx/mcp_tool.go b/internal/pkg/toolx/mcp_tool.go new file mode 100644 index 0000000..1742a1c --- /dev/null +++ b/internal/pkg/toolx/mcp_tool.go @@ -0,0 +1,86 @@ +package toolx + +import ( + "encoding/json" + "strings" + + "cs-agent/internal/pkg/dto/request" + "cs-agent/internal/pkg/errorsx" +) + +func BuildMCPToolCode(serverCode, toolName string) string { + serverCode = strings.TrimSpace(serverCode) + toolName = strings.TrimSpace(toolName) + if serverCode == "" || toolName == "" { + return "" + } + return serverCode + "/" + toolName +} + +func SplitMCPToolCode(toolCode string) (string, string) { + toolCode = strings.TrimSpace(toolCode) + if toolCode == "" { + return "", "" + } + idx := strings.Index(toolCode, "/") + if idx <= 0 || idx >= len(toolCode)-1 { + return "", "" + } + return strings.TrimSpace(toolCode[:idx]), strings.TrimSpace(toolCode[idx+1:]) +} + +func NormalizeMCPToolRequest(item request.AIAgentMCPToolRequest) (request.AIAgentMCPToolRequest, error) { + toolCode := strings.TrimSpace(item.ToolCode) + serverCode := strings.TrimSpace(item.ServerCode) + toolName := strings.TrimSpace(item.ToolName) + if toolCode != "" { + parsedServerCode, parsedToolName := SplitMCPToolCode(toolCode) + if parsedServerCode == "" || parsedToolName == "" { + return request.AIAgentMCPToolRequest{}, errorsx.InvalidParam("Direct Tool 的 toolCode 格式不合法") + } + if serverCode != "" && !strings.EqualFold(serverCode, parsedServerCode) { + return request.AIAgentMCPToolRequest{}, errorsx.InvalidParam("Direct Tool 的 toolCode 与 serverCode 不一致") + } + if toolName != "" && !strings.EqualFold(toolName, parsedToolName) { + return request.AIAgentMCPToolRequest{}, errorsx.InvalidParam("Direct Tool 的 toolCode 与 toolName 不一致") + } + serverCode = parsedServerCode + toolName = parsedToolName + } else { + toolCode = BuildMCPToolCode(serverCode, toolName) + } + if toolCode == "" || serverCode == "" || toolName == "" { + return request.AIAgentMCPToolRequest{}, errorsx.InvalidParam("Direct Tool 的 toolCode、serverCode 和 toolName 不能为空") + } + ret := request.AIAgentMCPToolRequest{ + ToolCode: toolCode, + ServerCode: serverCode, + ToolName: toolName, + Title: strings.TrimSpace(item.Title), + Description: strings.TrimSpace(item.Description), + } + if len(item.Arguments) > 0 { + ret.Arguments = make(map[string]string, len(item.Arguments)) + for key, value := range item.Arguments { + key = strings.TrimSpace(key) + value = strings.TrimSpace(value) + if key == "" || value == "" { + continue + } + ret.Arguments[key] = value + } + } + return ret, nil +} + +func ParseAgentMCPToolsJSON(raw string) ([]request.AIAgentMCPToolRequest, error) { + raw = strings.TrimSpace(raw) + if raw == "" { + return nil, nil + } + var ret []request.AIAgentMCPToolRequest + if err := json.Unmarshal([]byte(raw), &ret); err != nil { + return nil, err + } + return ret, nil +} diff --git a/internal/services/ai_agent_service.go b/internal/services/ai_agent_service.go index 231fc58..f5bddc0 100644 --- a/internal/services/ai_agent_service.go +++ b/internal/services/ai_agent_service.go @@ -12,6 +12,7 @@ import ( "cs-agent/internal/pkg/dto/request" "cs-agent/internal/pkg/enums" "cs-agent/internal/pkg/errorsx" + "cs-agent/internal/pkg/toolx" "cs-agent/internal/pkg/utils" "cs-agent/internal/repositories" @@ -294,37 +295,20 @@ func (s *aIAgentService) normalizeDirectTools(input []request.AIAgentMCPToolRequ ret := make([]request.AIAgentMCPToolRequest, 0, len(input)) seen := make(map[string]struct{}) for _, item := range input { - serverCode := strings.TrimSpace(item.ServerCode) - toolName := strings.TrimSpace(item.ToolName) - if serverCode == "" || toolName == "" { - return nil, errorsx.InvalidParam("Direct Tool 的 serverCode 和 toolName 不能为空") + normalized, err := toolx.NormalizeMCPToolRequest(item) + if err != nil { + return nil, err } + serverCode := strings.TrimSpace(normalized.ServerCode) server, ok := cfg.MCP.Servers[serverCode] if !ok || !server.Enabled { return nil, errorsx.InvalidParam("Direct Tool 绑定的 MCP 服务不存在或未启用") } - key := serverCode + "/" + toolName + key := strings.TrimSpace(normalized.ToolCode) if _, exists := seen[key]; exists { continue } seen[key] = struct{}{} - normalized := request.AIAgentMCPToolRequest{ - ServerCode: serverCode, - ToolName: toolName, - Title: strings.TrimSpace(item.Title), - Description: strings.TrimSpace(item.Description), - } - if len(item.Arguments) > 0 { - normalized.Arguments = make(map[string]string, len(item.Arguments)) - for key, value := range item.Arguments { - key = strings.TrimSpace(key) - value = strings.TrimSpace(value) - if key == "" || value == "" { - continue - } - normalized.Arguments[key] = value - } - } ret = append(ret, normalized) } return ret, nil