feat: add tool search functionality and enhance tool management in AI agent

This commit is contained in:
mlogclub
2026-04-10 14:57:46 +08:00
parent 1370e4b675
commit 8961ec4342
7 changed files with 435 additions and 33 deletions
@@ -109,7 +109,7 @@ func (s *Service) Run(ctx context.Context, req Request) (*Summary, error) {
collector.Data.Skill.RouteReason = summary.SkillRouteReason
collector.Data.Skill.RouteTrace = summary.SkillRouteTrace
agent, err := s.agentFactory.BuildCustomerServiceAgent(ctx, req.AIAgent, req.AIConfig, req.SelectedSkill, filteredToolDefs, req.ExtraTools, req.ExtraToolCodes, collector)
agent, err := s.agentFactory.BuildCustomerServiceAgent(ctx, req.AIAgent, req.AIConfig, req.SelectedSkill, filteredToolDefs, nil, req.ExtraTools, req.ExtraToolCodes, collector)
if err != nil {
summary.Status = "error"
summary.ErrorMessage = err.Error()
@@ -233,7 +233,7 @@ func (s *Service) Resume(ctx context.Context, req ResumeRequest) (*Summary, erro
collector.Data.Input.ToolCodes = append(collector.Data.Input.ToolCodes, summary.ToolCodes...)
collector.Data.Model.Provider = string(req.AIConfig.Provider)
collector.Data.Model.Name = req.AIConfig.ModelName
agent, err := s.agentFactory.BuildCustomerServiceAgent(ctx, req.AIAgent, req.AIConfig, nil, toolDefs, req.ExtraTools, req.ExtraToolCodes, collector)
agent, err := s.agentFactory.BuildCustomerServiceAgent(ctx, req.AIAgent, req.AIConfig, nil, nil, nil, req.ExtraTools, req.ExtraToolCodes, collector)
if err != nil {
summary.Status = "error"
summary.ErrorMessage = err.Error()
@@ -11,6 +11,7 @@ import (
einoagents "cs-agent/internal/ai/runtime/internal/impl/agents"
einocallbacks "cs-agent/internal/ai/runtime/internal/impl/callbacks"
"cs-agent/internal/models"
"cs-agent/internal/pkg/toolx"
"github.com/cloudwego/eino/adk"
einobasetool "github.com/cloudwego/eino/components/tool"
@@ -32,7 +33,8 @@ func NewAgentFactory() *AgentFactory {
}
func (f *AgentFactory) BuildCustomerServiceAgent(ctx context.Context, aiAgent *models.AIAgent, aiConfig *models.AIConfig,
selectedSkill *models.SkillDefinition, toolDefinitions []einoadapter.MCPToolDefinition, extraTools []einobasetool.BaseTool, extraToolCodes map[string]string,
selectedSkill *models.SkillDefinition, instructionToolDefinitions []einoadapter.MCPToolDefinition, mcpToolDefinitions []einoadapter.MCPToolDefinition,
extraTools []einobasetool.BaseTool, extraToolCodes map[string]string,
collector *einocallbacks.RuntimeTraceCollector) (*einoagents.CustomerServiceAgent, error) {
if aiAgent == nil || aiConfig == nil {
return nil, nil
@@ -41,7 +43,7 @@ func (f *AgentFactory) BuildCustomerServiceAgent(ctx context.Context, aiAgent *m
if err != nil {
return nil, err
}
baseTools, err := f.toolFactory.BuildBaseToolsByDefinitions(ctx, toolDefinitions)
baseTools, err := f.toolFactory.BuildBaseToolsByDefinitions(ctx, mcpToolDefinitions)
if err != nil {
return nil, err
}
@@ -50,20 +52,40 @@ func (f *AgentFactory) BuildCustomerServiceAgent(ctx context.Context, aiAgent *m
allTools = append(allTools, baseTools...)
handlers := make([]adk.ChatModelAgentMiddleware, 0, 1)
if collector != nil {
toolMetadataBy := make(map[string]einocallbacks.ToolMetadata, len(toolDefinitions))
for _, item := range toolDefinitions {
toolMetadataBy := make(map[string]einocallbacks.ToolMetadata, len(mcpToolDefinitions)+len(extraToolCodes))
for _, item := range mcpToolDefinitions {
toolMetadataBy[item.ModelName] = einocallbacks.ToolMetadata{
ToolCode: item.ToolCode,
ServerCode: item.ServerCode,
ToolName: item.ToolName,
}
}
for modelName, toolCode := range extraToolCodes {
modelName = strings.TrimSpace(modelName)
toolCode = strings.TrimSpace(toolCode)
if modelName == "" || toolCode == "" {
continue
}
serverCode, toolName := "", ""
if toolCode == toolx.BuiltinToolSearchToolCode {
serverCode = toolx.BuiltinToolCatalogServerCode
toolName = toolx.BuiltinToolSearchToolName
} else if toolCode == toolx.BuiltinCreateTicketConfirmToolCode {
serverCode = toolx.BuiltinToolCatalogServerCode
toolName = toolx.BuiltinCreateTicketConfirmToolName
}
toolMetadataBy[modelName] = einocallbacks.ToolMetadata{
ToolCode: toolCode,
ServerCode: serverCode,
ToolName: toolName,
}
}
handlers = append(handlers, einocallbacks.NewRuntimeTraceHandler(collector, toolMetadataBy))
}
inner, err := adk.NewChatModelAgent(ctx, &adk.ChatModelAgentConfig{
Name: strings.TrimSpace(aiAgent.Name),
Description: strings.TrimSpace(aiAgent.Description),
Instruction: buildAgentInstruction(aiAgent, selectedSkill, toolDefinitions, extraToolCodes),
Instruction: buildAgentInstruction(aiAgent, selectedSkill, instructionToolDefinitions, extraToolCodes),
Model: chatModel,
ToolsConfig: adk.ToolsConfig{
ToolsNodeConfig: compose.ToolsNodeConfig{
@@ -87,7 +109,16 @@ func buildAgentInstruction(aiAgent *models.AIAgent, selectedSkill *models.SkillD
if skillInstruction := buildSelectedSkillInstruction(selectedSkill, toolDefinitions); skillInstruction != "" {
appendixParts = append(appendixParts, skillInstruction)
}
if hasToolCode(extraToolCodes, "builtin/create_ticket_with_confirmation") {
if hasToolCode(extraToolCodes, toolx.BuiltinToolSearchToolCode) {
appendixParts = append(appendixParts, strings.TrimSpace(`
当你需要使用长尾 MCP 能力时,优先使用 tool_search 工具,并遵守以下规则:
1. 先用 query 搜索候选工具,再根据返回的 toolCode 选择真正的目标工具。
2. 只有在你已经明确要调用哪个动态工具时,才传入 toolCode 和 arguments 进行执行。
3. 不要臆造 toolCode;必须以 tool_search 返回的候选结果为准。
4. 如果当前已有固定内置工具可以完成任务,优先使用固定工具,不要滥用 tool_search。
`))
}
if hasToolCode(extraToolCodes, toolx.BuiltinCreateTicketConfirmToolCode) {
appendixParts = append(appendixParts, strings.TrimSpace(`
你可以在确认信息充分后调用 create_ticket_with_confirmation 工具来创建工单,但必须遵守以下规则:
1. 只有在用户明确表达希望提交工单、投诉、报障、售后处理等诉求时,才考虑调用该工具。
+39 -4
View File
@@ -19,6 +19,7 @@ func newService() *service {
return &service{
runtime: engine.NewService(),
registry: registry.NewRegistry(
tools.NewToolSearchTool(),
tools.NewCreateTicketConfirmTool(),
),
}
@@ -94,7 +95,7 @@ func (s *service) prepareToolsForRun(req *Request) error {
AIAgent: req.AIAgent,
AIConfig: req.AIConfig,
UserMessage: req.UserMessage,
AllowedToolCodes: resolveAllowedToolCodes(req.AIAgent, req.SelectedSkill),
AllowedToolCodes: expandRuntimeAllowedToolCodes(resolveAllowedToolCodes(req.AIAgent, req.SelectedSkill)),
})
if err != nil {
return err
@@ -109,9 +110,10 @@ func (s *service) prepareToolsForResume(req *ResumeRequest) error {
return nil
}
toolSet, err := s.registry.Resolve(registry.Context{
Conversation: req.Conversation,
AIAgent: req.AIAgent,
AIConfig: req.AIConfig,
Conversation: req.Conversation,
AIAgent: req.AIAgent,
AIConfig: req.AIConfig,
AllowedToolCodes: expandRuntimeAllowedToolCodes(parseAgentAllowedToolCodes(req.AIAgent)),
})
if err != nil {
return err
@@ -275,3 +277,36 @@ func resolveAllowedToolCodes(aiAgent *models.AIAgent, skill *models.SkillDefinit
return ret
}
}
func expandRuntimeAllowedToolCodes(items []string) []string {
ret := make([]string, 0, len(items)+1)
hasMCPTool := false
for _, item := range items {
item = strings.TrimSpace(item)
if item == "" {
continue
}
ret = append(ret, item)
serverCode, toolName := toolx.SplitMCPToolCode(item)
if serverCode != "" && toolName != "" {
hasMCPTool = true
}
}
if hasMCPTool {
ret = appendIfMissingString(ret, toolx.BuiltinToolSearchToolCode)
}
return ret
}
func appendIfMissingString(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)
}
@@ -0,0 +1,313 @@
package tools
import (
"context"
"encoding/json"
"fmt"
"slices"
"strings"
"cs-agent/internal/ai/mcps"
"cs-agent/internal/ai/runtime/registry"
"cs-agent/internal/pkg/toolx"
einotool "github.com/cloudwego/eino/components/tool"
"github.com/cloudwego/eino/schema"
einojsonschema "github.com/eino-contrib/jsonschema"
orderedmap "github.com/wk8/go-ordered-map/v2"
)
const (
ToolSearchToolCode = toolx.BuiltinToolSearchToolCode
ToolSearchToolName = toolx.BuiltinToolSearchToolName
)
type ToolSearchTool struct {
allowedToolCodes []string
}
func NewToolSearchTool() *ToolSearchTool {
return &ToolSearchTool{}
}
func (t *ToolSearchTool) Name() string {
return ToolSearchToolName
}
func (t *ToolSearchTool) Code() string {
return ToolSearchToolCode
}
func (t *ToolSearchTool) Enabled(ctx registry.Context) bool {
return len(filterAllowedMCPToolCodes(ctx.AllowedToolCodes)) > 0
}
func (t *ToolSearchTool) Build(ctx registry.Context) (einotool.BaseTool, error) {
if !t.Enabled(ctx) {
return nil, nil
}
return &ToolSearchTool{
allowedToolCodes: filterAllowedMCPToolCodes(ctx.AllowedToolCodes),
}, nil
}
func (t *ToolSearchTool) Info(ctx context.Context) (*schema.ToolInfo, error) {
return &schema.ToolInfo{
Name: ToolSearchToolName,
Desc: "当你需要使用当前会话允许的长尾 MCP 工具时,先调用本工具搜索合适的 toolCode;确认目标后,可再次调用本工具并传入 toolCode 与 arguments 代理执行。不要用它替代明确固定的内置流程工具。",
ParamsOneOf: schema.NewParamsOneOfByJSONSchema(&einojsonschema.Schema{
Version: einojsonschema.Version,
Type: "object",
Properties: orderedmap.New[string, *einojsonschema.Schema](orderedmap.WithInitialData(
orderedmap.Pair[string, *einojsonschema.Schema]{
Key: "query",
Value: &einojsonschema.Schema{
Type: "string",
Description: "要搜索的工具意图、能力或关键词;当只想列出候选工具时使用。",
},
},
orderedmap.Pair[string, *einojsonschema.Schema]{
Key: "toolCode",
Value: &einojsonschema.Schema{
Type: "string",
Description: "已确定目标后要调用的 MCP toolCode,例如 mcp_server/tool_name。",
},
},
orderedmap.Pair[string, *einojsonschema.Schema]{
Key: "arguments",
Value: &einojsonschema.Schema{
Type: "object",
Description: "调用目标工具时传入的参数对象。",
AdditionalProperties: &einojsonschema.Schema{},
},
},
)),
}),
Extra: map[string]any{
"toolCode": ToolSearchToolCode,
},
}, nil
}
func (t *ToolSearchTool) InvokableRun(ctx context.Context, argumentsInJSON string, opts ...einotool.Option) (string, error) {
if t == nil {
return "", fmt.Errorf("tool search tool is nil")
}
req, err := parseToolSearchRequest(argumentsInJSON)
if err != nil {
return "", err
}
if req.ToolCode != "" {
return t.invokeTargetTool(ctx, req.ToolCode, req.Arguments)
}
return t.searchCandidates(ctx, req.Query)
}
type toolSearchRequest struct {
Query string
ToolCode string
Arguments map[string]any
}
type toolSearchCandidate struct {
ToolCode string `json:"toolCode"`
ServerCode string `json:"serverCode"`
ToolName string `json:"toolName"`
Title string `json:"title,omitempty"`
Description string `json:"description,omitempty"`
}
func parseToolSearchRequest(argumentsInJSON string) (*toolSearchRequest, error) {
argumentsInJSON = strings.TrimSpace(argumentsInJSON)
if argumentsInJSON == "" {
return &toolSearchRequest{}, nil
}
raw := make(map[string]any)
if err := json.Unmarshal([]byte(argumentsInJSON), &raw); err != nil {
return nil, fmt.Errorf("invalid tool_search arguments: %w", err)
}
req := &toolSearchRequest{
Query: strings.TrimSpace(getStringValue(raw, "query")),
ToolCode: strings.TrimSpace(getStringValue(raw, "toolCode")),
}
if value, ok := raw["arguments"]; ok {
args, ok := value.(map[string]any)
if !ok {
return nil, fmt.Errorf("tool_search arguments must be an object")
}
req.Arguments = args
}
return req, nil
}
func (t *ToolSearchTool) searchCandidates(ctx context.Context, query string) (string, error) {
candidates, err := t.loadAllowedCandidates(ctx)
if err != nil {
return "", err
}
matched := filterCandidatesByQuery(candidates, query)
if len(matched) == 0 {
return "未找到匹配的动态工具,请换个关键词,或继续向用户追问后再搜索。", nil
}
if len(matched) > 8 {
matched = matched[:8]
}
buf, err := json.Marshal(map[string]any{
"query": strings.TrimSpace(query),
"total": len(matched),
"candidates": matched,
})
if err != nil {
return "", err
}
return string(buf), nil
}
func (t *ToolSearchTool) invokeTargetTool(ctx context.Context, toolCode string, arguments map[string]any) (string, error) {
toolCode = strings.TrimSpace(toolCode)
serverCode, toolName := toolx.SplitMCPToolCode(toolCode)
if serverCode == "" || toolName == "" {
return "", fmt.Errorf("tool_search 只支持调用 MCP toolCode")
}
if !containsToolCode(t.allowedToolCodes, toolCode) {
return "", fmt.Errorf("目标工具未被当前会话授权")
}
result, err := mcps.Runtime.CallTool(ctx, serverCode, toolName, cloneArguments(arguments))
if err != nil {
return "", err
}
return buildToolCallResultSummary(result), nil
}
func (t *ToolSearchTool) loadAllowedCandidates(ctx context.Context) ([]toolSearchCandidate, error) {
serverToToolCodes := make(map[string]map[string]struct{})
for _, toolCode := range t.allowedToolCodes {
serverCode, toolName := toolx.SplitMCPToolCode(toolCode)
if serverCode == "" || toolName == "" {
continue
}
if _, ok := serverToToolCodes[serverCode]; !ok {
serverToToolCodes[serverCode] = make(map[string]struct{})
}
serverToToolCodes[serverCode][toolCode] = struct{}{}
}
serverCodes := make([]string, 0, len(serverToToolCodes))
for serverCode := range serverToToolCodes {
serverCodes = append(serverCodes, serverCode)
}
slices.Sort(serverCodes)
ret := make([]toolSearchCandidate, 0)
for _, serverCode := range serverCodes {
tools, err := mcps.Runtime.ListTools(ctx, serverCode)
if err != nil {
return nil, err
}
allowed := serverToToolCodes[serverCode]
for _, item := range tools {
toolCode := toolx.BuildMCPToolCode(serverCode, item.Name)
if _, ok := allowed[toolCode]; !ok {
continue
}
ret = append(ret, toolSearchCandidate{
ToolCode: toolCode,
ServerCode: serverCode,
ToolName: strings.TrimSpace(item.Name),
Title: strings.TrimSpace(item.Title),
Description: strings.TrimSpace(item.Description),
})
}
}
return ret, nil
}
func filterAllowedMCPToolCodes(input []string) []string {
if len(input) == 0 {
return nil
}
ret := make([]string, 0, len(input))
for _, item := range input {
item = strings.TrimSpace(item)
serverCode, toolName := toolx.SplitMCPToolCode(item)
if serverCode == "" || toolName == "" {
continue
}
ret = append(ret, item)
}
return ret
}
func containsToolCode(items []string, target string) bool {
target = strings.TrimSpace(target)
if target == "" {
return false
}
for _, item := range items {
if strings.TrimSpace(item) == target {
return true
}
}
return false
}
func filterCandidatesByQuery(candidates []toolSearchCandidate, query string) []toolSearchCandidate {
query = strings.TrimSpace(strings.ToLower(query))
if query == "" {
return candidates
}
ret := make([]toolSearchCandidate, 0, len(candidates))
for _, item := range candidates {
searchText := strings.ToLower(strings.Join([]string{
item.ToolCode,
item.ServerCode,
item.ToolName,
item.Title,
item.Description,
}, "\n"))
if strings.Contains(searchText, query) {
ret = append(ret, item)
}
}
return ret
}
func cloneArguments(input map[string]any) map[string]any {
if len(input) == 0 {
return map[string]any{}
}
ret := make(map[string]any, len(input))
for key, value := range input {
ret[key] = value
}
return ret
}
func buildToolCallResultSummary(result *mcps.ToolCallResult) string {
if result == nil {
return ""
}
lines := make([]string, 0, len(result.Content)+2)
if result.IsError {
lines = append(lines, "tool returned an error")
}
if result.StructuredContent != nil {
if data, err := json.Marshal(result.StructuredContent); err == nil {
lines = append(lines, string(data))
}
}
for _, item := range result.Content {
switch item.Type {
case "text":
if text := strings.TrimSpace(item.Text); text != "" {
lines = append(lines, text)
}
default:
if item.Data == nil {
continue
}
if data, err := json.Marshal(item.Data); err == nil {
lines = append(lines, string(data))
}
}
}
return strings.TrimSpace(strings.Join(lines, "\n"))
}
@@ -176,7 +176,10 @@ func buildAIAgentResponse(item *models.AIAgent) response.AIAgentResponse {
}
serverCode := strings.TrimSpace(tool.ServerCode)
toolName := strings.TrimSpace(tool.ToolName)
if toolCode == toolx.BuiltinCreateTicketConfirmToolCode {
if toolCode == toolx.BuiltinToolSearchToolCode {
serverCode = toolx.BuiltinToolCatalogServerCode
toolName = toolx.BuiltinToolSearchToolName
} else if toolCode == toolx.BuiltinCreateTicketConfirmToolCode {
serverCode = toolx.BuiltinToolCatalogServerCode
toolName = toolx.BuiltinCreateTicketConfirmToolName
} else if parsedServerCode, parsedToolName := toolx.SplitMCPToolCode(toolCode); parsedServerCode != "" && parsedToolName != "" {
@@ -184,12 +187,22 @@ func buildAIAgentResponse(item *models.AIAgent) response.AIAgentResponse {
toolName = parsedToolName
}
title := strings.TrimSpace(tool.Title)
if title == "" && toolCode == toolx.BuiltinCreateTicketConfirmToolCode {
title = toolx.BuiltinCreateTicketConfirmToolTitle
if title == "" {
switch toolCode {
case toolx.BuiltinToolSearchToolCode:
title = toolx.BuiltinToolSearchToolTitle
case toolx.BuiltinCreateTicketConfirmToolCode:
title = toolx.BuiltinCreateTicketConfirmToolTitle
}
}
description := strings.TrimSpace(tool.Description)
if description == "" && toolCode == toolx.BuiltinCreateTicketConfirmToolCode {
description = toolx.BuiltinCreateTicketConfirmToolDescription
if description == "" {
switch toolCode {
case toolx.BuiltinToolSearchToolCode:
description = toolx.BuiltinToolSearchToolDescription
case toolx.BuiltinCreateTicketConfirmToolCode:
description = toolx.BuiltinCreateTicketConfirmToolDescription
}
}
ret.DirectTools = append(ret.DirectTools, response.AIAgentMCPToolResponse{
ToolCode: toolCode,
+4
View File
@@ -2,6 +2,10 @@ package toolx
const (
BuiltinToolCatalogServerCode = "builtin"
BuiltinToolSearchToolCode = "builtin/tool_search"
BuiltinToolSearchToolName = "tool_search"
BuiltinToolSearchToolTitle = "搜索并调用动态工具"
BuiltinToolSearchToolDescription = "用于搜索当前允许使用的 MCP 工具,并在确认目标 toolCode 后动态调用该工具。适合处理长尾工具,不应替代固定内置流程工具。"
BuiltinCreateTicketConfirmToolCode = "builtin/create_ticket_with_confirmation"
BuiltinCreateTicketConfirmToolName = "create_ticket_with_confirmation"
BuiltinCreateTicketConfirmToolTitle = "创建工单并发起确认"
+22 -16
View File
@@ -32,21 +32,7 @@ type MCPToolCatalogItem struct {
func (s *toolCatalogService) ListMCPTools(ctx context.Context) ([]MCPToolCatalogItem, error) {
cfg := config.Current()
if !cfg.MCP.Enabled {
return nil, errorsx.InvalidParam("MCP未启用")
}
if len(cfg.MCP.Servers) == 0 {
return nil, nil
}
serverCodes := make([]string, 0, len(cfg.MCP.Servers))
for serverCode, server := range cfg.MCP.Servers {
if !server.Enabled {
continue
}
serverCodes = append(serverCodes, serverCode)
}
slices.Sort(serverCodes)
ret := make([]MCPToolCatalogItem, 0)
ret := make([]MCPToolCatalogItem, 0, 2)
ret = append(ret, MCPToolCatalogItem{
ToolCode: toolx.BuiltinCreateTicketConfirmToolCode,
ServerCode: toolx.BuiltinToolCatalogServerCode,
@@ -55,6 +41,25 @@ func (s *toolCatalogService) ListMCPTools(ctx context.Context) ([]MCPToolCatalog
Title: toolx.BuiltinCreateTicketConfirmToolTitle,
Description: toolx.BuiltinCreateTicketConfirmToolDescription,
})
if !cfg.MCP.Enabled {
return ret, nil
}
ret = append(ret, MCPToolCatalogItem{
ToolCode: toolx.BuiltinToolSearchToolCode,
ServerCode: toolx.BuiltinToolCatalogServerCode,
ToolName: toolx.BuiltinToolSearchToolName,
SourceType: toolx.BuiltinToolCatalogServerCode,
Title: toolx.BuiltinToolSearchToolTitle,
Description: toolx.BuiltinToolSearchToolDescription,
})
serverCodes := make([]string, 0, len(cfg.MCP.Servers))
for serverCode, server := range cfg.MCP.Servers {
if !server.Enabled {
continue
}
serverCodes = append(serverCodes, serverCode)
}
slices.Sort(serverCodes)
for _, serverCode := range serverCodes {
tools, err := mcps.Runtime.ListTools(ctx, serverCode)
if err != nil {
@@ -86,7 +91,8 @@ func (s *toolCatalogService) ValidateToolCode(toolCode string) error {
if toolCode == "" {
return errorsx.InvalidParam("toolCode不能为空")
}
if toolCode == toolx.BuiltinCreateTicketConfirmToolCode {
switch toolCode {
case toolx.BuiltinToolSearchToolCode, toolx.BuiltinCreateTicketConfirmToolCode:
return nil
}
serverCode, toolName := toolx.SplitMCPToolCode(toolCode)