diff --git a/internal/ai/runtime/executor/service.go b/internal/ai/runtime/executor/service.go index 5e2b824..79cb13e 100644 --- a/internal/ai/runtime/executor/service.go +++ b/internal/ai/runtime/executor/service.go @@ -214,4 +214,7 @@ func syncSkillSummaryFromCollector(summary *RunResult, collector *callbacks.Runt summary.SkillRouteReason = strings.TrimSpace(trace.RouteReason) summary.SkillRouteTrace = strings.TrimSpace(trace.RouteTrace) summary.SkillAllowedToolCodes = append([]string(nil), trace.AllowedToolCodes...) + if len(trace.FilteredToolCodes) > 0 { + summary.ToolCodes = append([]string(nil), trace.FilteredToolCodes...) + } } diff --git a/internal/ai/runtime/internal/impl/callbacks/runlog_callback.go b/internal/ai/runtime/internal/impl/callbacks/runlog_callback.go index 7d628b0..9aeb47e 100644 --- a/internal/ai/runtime/internal/impl/callbacks/runlog_callback.go +++ b/internal/ai/runtime/internal/impl/callbacks/runlog_callback.go @@ -100,6 +100,15 @@ func (c *RuntimeTraceCollector) ActivateSkill(skill SkillMetadata, routeReason s c.Data.Skill.RouteTrace = routeTrace } +func (c *RuntimeTraceCollector) SetFilteredToolCodes(toolCodes []string) { + if c == nil { + return + } + c.mu.Lock() + defer c.mu.Unlock() + c.Data.Skill.FilteredToolCodes = append([]string(nil), toolCodes...) +} + func (c *RuntimeTraceCollector) SetRetrieverSummary(summary RetrieverTraceSummary) { if c == nil { return diff --git a/internal/ai/runtime/internal/impl/callbacks/trace_callback.go b/internal/ai/runtime/internal/impl/callbacks/trace_callback.go index e070c2d..660f1bc 100644 --- a/internal/ai/runtime/internal/impl/callbacks/trace_callback.go +++ b/internal/ai/runtime/internal/impl/callbacks/trace_callback.go @@ -12,6 +12,8 @@ type ToolTraceItem struct { LatencyMs int64 `json:"latencyMs,omitempty"` Status string `json:"status,omitempty"` ErrorMessage string `json:"errorMessage,omitempty"` + Blocked bool `json:"blocked,omitempty"` + BlockedReason string `json:"blockedReason,omitempty"` } type ToolSearchTraceItem struct { @@ -146,6 +148,7 @@ type SkillTraceData struct { RouteReason string `json:"routeReason,omitempty"` RouteTrace string `json:"routeTrace,omitempty"` AllowedToolCodes []string `json:"allowedToolCodes,omitempty"` + FilteredToolCodes []string `json:"filteredToolCodes,omitempty"` MiddlewareEnabled bool `json:"middlewareEnabled,omitempty"` MiddlewareToolName string `json:"middlewareToolName,omitempty"` VisibleCodes []string `json:"visibleCodes,omitempty"` diff --git a/internal/ai/runtime/internal/impl/factory/agent_handler_service.go b/internal/ai/runtime/internal/impl/factory/agent_handler_service.go index 8ae25f8..5783b33 100644 --- a/internal/ai/runtime/internal/impl/factory/agent_handler_service.go +++ b/internal/ai/runtime/internal/impl/factory/agent_handler_service.go @@ -2,6 +2,7 @@ package factory import ( "context" + "strings" einocallbacks "cs-agent/internal/ai/runtime/internal/impl/callbacks" "cs-agent/internal/ai/runtime/registry" @@ -36,7 +37,26 @@ func NewAgentHandlerService(skillMiddleware *SkillMiddlewareService) *AgentHandl } func (s *AgentHandlerService) Build(ctx context.Context, input BuildAgentHandlersInput) ([]adk.ChatModelAgentMiddleware, error) { - handlers := make([]adk.ChatModelAgentMiddleware, 0, 3) + handlers := make([]adk.ChatModelAgentMiddleware, 0, 4) + skillMetadataByCode := buildRuntimeSkillMetadataMap(input.AIAgent) + toolMetadataBy := buildRuntimeTraceToolMetadata(input.DynamicToolDefinitions, input.StaticToolMetadata, len(skillMetadataByCode) > 0) + traceSkillMetadata := make(map[string]einocallbacks.SkillMetadata, len(skillMetadataByCode)) + for code, item := range skillMetadataByCode { + traceSkillMetadata[code] = einocallbacks.SkillMetadata{ + Code: item.Code, + Name: item.Name, + Description: item.Description, + AllowedToolCodes: append([]string(nil), item.AllowedToolCodes...), + } + } + if input.Collector != nil { + if len(skillMetadataByCode) > 0 { + input.Collector.SetSkillMiddleware(true, toolx.BuiltinSkill.Name) + } + input.Collector.SetVisibleSkills(traceSkillMetadata) + input.Collector.SetInstructionSummary(input.InstructionSummary) + handlers = append(handlers, einocallbacks.NewRuntimeTraceHandler(input.Collector, toolMetadataBy, traceSkillMetadata)) + } if len(input.DynamicTools) > 0 { toolSearchHandler, err := einotoolsearch.New(ctx, &einotoolsearch.Config{ DynamicTools: input.DynamicTools, @@ -46,31 +66,33 @@ func (s *AgentHandlerService) Build(ctx context.Context, input BuildAgentHandler } handlers = append(handlers, toolSearchHandler) } - skillMetadataByCode := buildRuntimeSkillMetadataMap(input.AIAgent) if len(skillMetadataByCode) > 0 { skillHandler, err := s.skillMiddleware.Build(ctx, input.AIAgent, input.InstructionToolDefinitions) if err != nil { return nil, err } handlers = append(handlers, skillHandler) - } - if input.Collector != nil { - toolMetadataBy := buildRuntimeTraceToolMetadata(input.DynamicToolDefinitions, input.StaticToolMetadata, len(skillMetadataByCode) > 0) - traceSkillMetadata := make(map[string]einocallbacks.SkillMetadata, len(skillMetadataByCode)) - for code, item := range skillMetadataByCode { - traceSkillMetadata[code] = einocallbacks.SkillMetadata{ - Code: item.Code, - Name: item.Name, - Description: item.Description, - AllowedToolCodes: append([]string(nil), item.AllowedToolCodes...), - } - } - if len(skillMetadataByCode) > 0 { - input.Collector.SetSkillMiddleware(true, toolx.BuiltinSkill.Name) - } - input.Collector.SetVisibleSkills(traceSkillMetadata) - input.Collector.SetInstructionSummary(input.InstructionSummary) - handlers = append(handlers, einocallbacks.NewRuntimeTraceHandler(input.Collector, toolMetadataBy, traceSkillMetadata)) + handlers = append(handlers, NewRuntimeToolFilterMiddleware( + input.Collector, + toolMetadataBy, + traceSkillMetadata, + dynamicToolModelNames(input.DynamicToolDefinitions), + )) } return handlers, nil } + +func dynamicToolModelNames(definitions []runtimetooling.MCPToolDefinition) []string { + ret := make([]string, 0, len(definitions)) + for _, item := range definitions { + modelName := strings.TrimSpace(item.ModelName) + if modelName == "" { + modelName = strings.TrimSpace(runtimetooling.BuildModelToolName(item)) + } + if modelName == "" { + continue + } + ret = append(ret, modelName) + } + return ret +} diff --git a/internal/ai/runtime/internal/impl/factory/tool_filter_middleware.go b/internal/ai/runtime/internal/impl/factory/tool_filter_middleware.go new file mode 100644 index 0000000..bc63bd8 --- /dev/null +++ b/internal/ai/runtime/internal/impl/factory/tool_filter_middleware.go @@ -0,0 +1,402 @@ +package factory + +import ( + "context" + "encoding/json" + "fmt" + "strings" + + einocallbacks "cs-agent/internal/ai/runtime/internal/impl/callbacks" + "cs-agent/internal/pkg/toolx" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/components/model" + einotool "github.com/cloudwego/eino/components/tool" + "github.com/cloudwego/eino/schema" +) + +const activeSkillRunLocalKey = "runtime_active_skill_code" + +type RuntimeToolFilterMiddleware struct { + *adk.BaseChatModelAgentMiddleware + collector *einocallbacks.RuntimeTraceCollector + toolMetadataByName map[string]einocallbacks.ToolMetadata + skillMetadataBy map[string]einocallbacks.SkillMetadata + dynamicToolNames []string +} + +func NewRuntimeToolFilterMiddleware( + collector *einocallbacks.RuntimeTraceCollector, + toolMetadataByName map[string]einocallbacks.ToolMetadata, + skillMetadataBy map[string]einocallbacks.SkillMetadata, + dynamicToolNames []string, +) *RuntimeToolFilterMiddleware { + return &RuntimeToolFilterMiddleware{ + BaseChatModelAgentMiddleware: &adk.BaseChatModelAgentMiddleware{}, + collector: collector, + toolMetadataByName: cloneToolMetadataMap(toolMetadataByName), + skillMetadataBy: cloneSkillMetadataMap(skillMetadataBy), + dynamicToolNames: append([]string(nil), dynamicToolNames...), + } +} + +func (m *RuntimeToolFilterMiddleware) WrapModel(_ context.Context, cm model.BaseChatModel, mc *adk.ModelContext) (model.BaseChatModel, error) { + if mc == nil { + return cm, nil + } + return &runtimeToolFilterModelWrapper{ + cm: cm, + allTools: append([]*schema.ToolInfo(nil), mc.Tools...), + collector: m.collector, + toolMetadataByName: m.toolMetadataByName, + skillMetadataBy: m.skillMetadataBy, + dynamicToolNames: append([]string(nil), m.dynamicToolNames...), + }, nil +} + +func (m *RuntimeToolFilterMiddleware) WrapInvokableToolCall(_ context.Context, endpoint adk.InvokableToolCallEndpoint, tCtx *adk.ToolContext) (adk.InvokableToolCallEndpoint, error) { + return func(ctx context.Context, argumentsInJSON string, opts ...einotool.Option) (string, error) { + toolName := "" + if tCtx != nil { + toolName = strings.TrimSpace(tCtx.Name) + } + metadata, _ := resolveRuntimeToolMetadata(toolName, m.toolMetadataByName) + if !isRuntimeBuiltinAlwaysAllowed(metadata.ToolCode) { + activeSkill, restricted := m.resolveActiveSkill(ctx) + if restricted && !isToolCodeAllowedForSkill(metadata.ToolCode, activeSkill.AllowedToolCodes) { + return "", m.blockToolCall(metadata, argumentsInJSON, activeSkill) + } + } + result, err := endpoint(ctx, argumentsInJSON, opts...) + if err != nil { + return result, err + } + if strings.TrimSpace(metadata.ToolCode) == toolx.BuiltinSkill.Code { + _ = m.setActiveSkill(ctx, skillCodeFromArguments(argumentsInJSON)) + return result, nil + } + if strings.TrimSpace(metadata.ToolCode) == toolx.BuiltinToolSearch.Code { + activeSkill, restricted := m.resolveActiveSkill(ctx) + if !restricted { + return result, nil + } + filtered, filterErr := filterToolSearchResult(result, activeSkill.AllowedToolCodes, m.toolMetadataByName) + if filterErr == nil { + return filtered, nil + } + } + return result, nil + }, nil +} + +func (m *RuntimeToolFilterMiddleware) blockToolCall(metadata einocallbacks.ToolMetadata, argumentsInJSON string, activeSkill einocallbacks.SkillMetadata) error { + err := fmt.Errorf("tool %s is not allowed for active skill %s", strings.TrimSpace(metadata.ToolCode), strings.TrimSpace(activeSkill.Code)) + if m.collector != nil { + m.collector.AddToolItem(einocallbacks.ToolTraceItem{ + ToolCode: strings.TrimSpace(metadata.ToolCode), + ServerCode: strings.TrimSpace(metadata.ServerCode), + ToolName: strings.TrimSpace(metadata.ToolName), + Arguments: parseRuntimeToolArguments(argumentsInJSON), + Status: "error", + ErrorMessage: err.Error(), + Blocked: true, + BlockedReason: "skill_tool_not_allowed", + }) + } + return err +} + +func (m *RuntimeToolFilterMiddleware) setActiveSkill(ctx context.Context, skillCode string) error { + skillCode = strings.TrimSpace(skillCode) + if skillCode == "" { + return nil + } + return adk.SetRunLocalValue(ctx, activeSkillRunLocalKey, skillCode) +} + +func (m *RuntimeToolFilterMiddleware) resolveActiveSkill(ctx context.Context) (einocallbacks.SkillMetadata, bool) { + if len(m.skillMetadataBy) == 0 { + return einocallbacks.SkillMetadata{}, false + } + value, found, err := adk.GetRunLocalValue(ctx, activeSkillRunLocalKey) + if err != nil || !found { + return einocallbacks.SkillMetadata{}, false + } + code, ok := value.(string) + if !ok { + return einocallbacks.SkillMetadata{}, false + } + code = strings.TrimSpace(code) + if code == "" { + return einocallbacks.SkillMetadata{}, false + } + skill, ok := m.skillMetadataBy[code] + if !ok { + return einocallbacks.SkillMetadata{}, false + } + if len(skill.AllowedToolCodes) == 0 { + return skill, false + } + return skill, true +} + +type runtimeToolFilterModelWrapper struct { + cm model.BaseChatModel + allTools []*schema.ToolInfo + collector *einocallbacks.RuntimeTraceCollector + toolMetadataByName map[string]einocallbacks.ToolMetadata + skillMetadataBy map[string]einocallbacks.SkillMetadata + dynamicToolNames []string +} + +func (w *runtimeToolFilterModelWrapper) Generate(ctx context.Context, input []*schema.Message, opts ...model.Option) (*schema.Message, error) { + tools := w.filteredTools(ctx, input) + return w.cm.Generate(ctx, input, append(opts, model.WithTools(tools))...) +} + +func (w *runtimeToolFilterModelWrapper) Stream(ctx context.Context, input []*schema.Message, opts ...model.Option) (*schema.StreamReader[*schema.Message], error) { + tools := w.filteredTools(ctx, input) + return w.cm.Stream(ctx, input, append(opts, model.WithTools(tools))...) +} + +func (w *runtimeToolFilterModelWrapper) filteredTools(ctx context.Context, input []*schema.Message) []*schema.ToolInfo { + tools := filterDynamicToolInfos(w.allTools, w.dynamicToolNames, input) + activeSkill, restricted := resolveActiveSkillMetadata(ctx, w.skillMetadataBy) + if restricted { + tools = filterToolInfosBySkill(tools, w.toolMetadataByName, activeSkill.AllowedToolCodes) + } + if w.collector != nil { + w.collector.SetFilteredToolCodes(extractToolCodesFromInfos(tools, w.toolMetadataByName)) + } + return tools +} + +func resolveActiveSkillMetadata(ctx context.Context, skills map[string]einocallbacks.SkillMetadata) (einocallbacks.SkillMetadata, bool) { + if len(skills) == 0 { + return einocallbacks.SkillMetadata{}, false + } + value, found, err := adk.GetRunLocalValue(ctx, activeSkillRunLocalKey) + if err != nil || !found { + return einocallbacks.SkillMetadata{}, false + } + code, ok := value.(string) + if !ok { + return einocallbacks.SkillMetadata{}, false + } + code = strings.TrimSpace(code) + skill, ok := skills[code] + if !ok || len(skill.AllowedToolCodes) == 0 { + return skill, false + } + return skill, true +} + +func filterDynamicToolInfos(allTools []*schema.ToolInfo, dynamicToolNames []string, messages []*schema.Message) []*schema.ToolInfo { + if len(allTools) == 0 { + return nil + } + selectedToolNames := extractSelectedDynamicToolNames(messages) + if len(dynamicToolNames) == 0 { + return append([]*schema.ToolInfo(nil), allTools...) + } + removeMap := invertStringSelection(dynamicToolNames, selectedToolNames) + ret := make([]*schema.ToolInfo, 0, len(allTools)) + for _, info := range allTools { + if info == nil { + continue + } + if _, ok := removeMap[strings.TrimSpace(info.Name)]; ok { + continue + } + ret = append(ret, info) + } + return ret +} + +func extractSelectedDynamicToolNames(messages []*schema.Message) []string { + if len(messages) == 0 { + return nil + } + selected := make([]string, 0) + for _, message := range messages { + if message == nil || message.Role != schema.Tool || strings.TrimSpace(message.ToolName) != toolx.BuiltinToolSearch.Name { + continue + } + var payload struct { + SelectedTools []string `json:"selectedTools"` + } + if err := json.Unmarshal([]byte(strings.TrimSpace(message.Content)), &payload); err != nil { + continue + } + for _, item := range payload.SelectedTools { + item = strings.TrimSpace(item) + if item == "" { + continue + } + selected = append(selected, item) + } + } + return selected +} + +func invertStringSelection(all []string, selected []string) map[string]struct{} { + selectedSet := make(map[string]struct{}, len(selected)) + for _, item := range selected { + item = strings.TrimSpace(item) + if item == "" { + continue + } + selectedSet[item] = struct{}{} + } + ret := make(map[string]struct{}) + for _, item := range all { + item = strings.TrimSpace(item) + if item == "" { + continue + } + if _, ok := selectedSet[item]; ok { + continue + } + ret[item] = struct{}{} + } + return ret +} + +func filterToolInfosBySkill(allTools []*schema.ToolInfo, toolMetadataByName map[string]einocallbacks.ToolMetadata, allowedToolCodes []string) []*schema.ToolInfo { + if len(allTools) == 0 || len(allowedToolCodes) == 0 { + return append([]*schema.ToolInfo(nil), allTools...) + } + ret := make([]*schema.ToolInfo, 0, len(allTools)) + for _, info := range allTools { + if info == nil { + continue + } + metadata, _ := resolveRuntimeToolMetadata(strings.TrimSpace(info.Name), toolMetadataByName) + if isRuntimeBuiltinAlwaysAllowed(metadata.ToolCode) || isToolCodeAllowedForSkill(metadata.ToolCode, allowedToolCodes) { + ret = append(ret, info) + } + } + return ret +} + +func extractToolCodesFromInfos(infos []*schema.ToolInfo, toolMetadataByName map[string]einocallbacks.ToolMetadata) []string { + ret := make([]string, 0, len(infos)) + for _, info := range infos { + if info == nil { + continue + } + metadata, ok := resolveRuntimeToolMetadata(strings.TrimSpace(info.Name), toolMetadataByName) + if !ok || strings.TrimSpace(metadata.ToolCode) == "" { + continue + } + ret = append(ret, metadata.ToolCode) + } + return toolx.NormalizeToolCodes(ret) +} + +func isToolCodeAllowedForSkill(toolCode string, allowedToolCodes []string) bool { + toolCode = toolx.NormalizeToolCodeAlias(strings.TrimSpace(toolCode)) + if toolCode == "" { + return true + } + for _, item := range toolx.NormalizeToolCodes(allowedToolCodes) { + if item == toolCode { + return true + } + } + return false +} + +func isRuntimeBuiltinAlwaysAllowed(toolCode string) bool { + toolCode = toolx.NormalizeToolCodeAlias(strings.TrimSpace(toolCode)) + return toolCode == toolx.BuiltinSkill.Code || toolCode == toolx.BuiltinToolSearch.Code +} + +func resolveRuntimeToolMetadata(toolName string, toolMetadataByName map[string]einocallbacks.ToolMetadata) (einocallbacks.ToolMetadata, bool) { + toolName = strings.TrimSpace(toolName) + if toolName == "" { + return einocallbacks.ToolMetadata{}, false + } + if spec, ok := toolx.GetRegisteredToolSpecByName(toolName); ok { + resolved := toolx.ResolveToolMetadata(spec.Code, spec.Name) + return einocallbacks.ToolMetadata{ + ToolCode: resolved.ToolCode, + ServerCode: resolved.ServerCode, + ToolName: resolved.ToolName, + SourceType: resolved.SourceType, + }, true + } + metadata, ok := toolMetadataByName[toolName] + return metadata, ok +} + +func skillCodeFromArguments(argumentsInJSON string) string { + var args struct { + Skill string `json:"skill"` + } + if err := json.Unmarshal([]byte(strings.TrimSpace(argumentsInJSON)), &args); err != nil { + return "" + } + return strings.TrimSpace(args.Skill) +} + +func parseRuntimeToolArguments(argumentsInJSON string) map[string]any { + argumentsInJSON = strings.TrimSpace(argumentsInJSON) + if argumentsInJSON == "" { + return nil + } + ret := make(map[string]any) + if err := json.Unmarshal([]byte(argumentsInJSON), &ret); err != nil { + return nil + } + return ret +} + +func filterToolSearchResult(result string, allowedToolCodes []string, toolMetadataByName map[string]einocallbacks.ToolMetadata) (string, error) { + result = strings.TrimSpace(result) + if result == "" { + return result, nil + } + var payload struct { + SelectedTools []string `json:"selectedTools"` + } + if err := json.Unmarshal([]byte(result), &payload); err != nil { + return result, err + } + filtered := make([]string, 0, len(payload.SelectedTools)) + for _, toolName := range payload.SelectedTools { + metadata, _ := resolveRuntimeToolMetadata(toolName, toolMetadataByName) + if isToolCodeAllowedForSkill(metadata.ToolCode, allowedToolCodes) { + filtered = append(filtered, strings.TrimSpace(toolName)) + } + } + payload.SelectedTools = filtered + buf, err := json.Marshal(payload) + if err != nil { + return "", err + } + return string(buf), nil +} + +func cloneToolMetadataMap(input map[string]einocallbacks.ToolMetadata) map[string]einocallbacks.ToolMetadata { + if len(input) == 0 { + return nil + } + ret := make(map[string]einocallbacks.ToolMetadata, len(input)) + for key, value := range input { + ret[key] = value + } + return ret +} + +func cloneSkillMetadataMap(input map[string]einocallbacks.SkillMetadata) map[string]einocallbacks.SkillMetadata { + if len(input) == 0 { + return nil + } + ret := make(map[string]einocallbacks.SkillMetadata, len(input)) + for key, value := range input { + value.AllowedToolCodes = append([]string(nil), value.AllowedToolCodes...) + ret[key] = value + } + return ret +} diff --git a/internal/ai/runtime/internal/impl/factory/tool_filter_middleware_test.go b/internal/ai/runtime/internal/impl/factory/tool_filter_middleware_test.go new file mode 100644 index 0000000..7a86f6a --- /dev/null +++ b/internal/ai/runtime/internal/impl/factory/tool_filter_middleware_test.go @@ -0,0 +1,84 @@ +package factory + +import ( + "testing" + + einocallbacks "cs-agent/internal/ai/runtime/internal/impl/callbacks" + "cs-agent/internal/pkg/toolx" + + "github.com/cloudwego/eino/schema" +) + +func TestFilterDynamicToolInfos(t *testing.T) { + allTools := []*schema.ToolInfo{ + {Name: toolx.BuiltinToolSearch.Name}, + {Name: "mcp_server_a"}, + {Name: "mcp_server_b"}, + } + messages := []*schema.Message{ + {Role: schema.Tool, ToolName: toolx.BuiltinToolSearch.Name, Content: `{"selectedTools":["mcp_server_b"]}`}, + } + + filtered := filterDynamicToolInfos(allTools, []string{"mcp_server_a", "mcp_server_b"}, messages) + if len(filtered) != 2 { + t.Fatalf("unexpected filtered tool count: %d", len(filtered)) + } + if filtered[0].Name != toolx.BuiltinToolSearch.Name || filtered[1].Name != "mcp_server_b" { + t.Fatalf("unexpected filtered tools: %#v", filtered) + } +} + +func TestFilterToolInfosBySkill(t *testing.T) { + allTools := []*schema.ToolInfo{ + {Name: toolx.BuiltinSkill.Name}, + {Name: toolx.BuiltinToolSearch.Name}, + {Name: toolx.GraphHandoffConversation.Name}, + {Name: "mcp_server_refund"}, + } + toolMetadataByName := map[string]einocallbacks.ToolMetadata{ + toolx.GraphHandoffConversation.Name: {ToolCode: toolx.GraphHandoffConversation.Code, ToolName: toolx.GraphHandoffConversation.Name}, + "mcp_server_refund": {ToolCode: "mcp/refund", ToolName: "mcp_server_refund"}, + } + + filtered := filterToolInfosBySkill(allTools, toolMetadataByName, []string{toolx.GraphHandoffConversation.Code}) + if len(filtered) != 3 { + t.Fatalf("unexpected filtered tool count: %d", len(filtered)) + } + if filtered[0].Name != toolx.BuiltinSkill.Name || filtered[1].Name != toolx.BuiltinToolSearch.Name || filtered[2].Name != toolx.GraphHandoffConversation.Name { + t.Fatalf("unexpected filtered tools: %#v", filtered) + } +} + +func TestFilterToolSearchResult(t *testing.T) { + toolMetadataByName := map[string]einocallbacks.ToolMetadata{ + "mcp_server_refund": {ToolCode: "mcp/refund", ToolName: "mcp_server_refund"}, + "mcp_server_order": {ToolCode: "mcp/order", ToolName: "mcp_server_order"}, + } + + got, err := filterToolSearchResult(`{"selectedTools":["mcp_server_refund","mcp_server_order"]}`, []string{"mcp/order"}, toolMetadataByName) + if err != nil { + t.Fatalf("filterToolSearchResult returned error: %v", err) + } + if got != `{"selectedTools":["mcp_server_order"]}` { + t.Fatalf("unexpected filtered result: %s", got) + } +} + +func TestExtractToolCodesFromInfos(t *testing.T) { + infos := []*schema.ToolInfo{ + {Name: toolx.BuiltinSkill.Name}, + {Name: toolx.GraphPrepareTicketDraft.Name}, + {Name: "mcp_server_refund"}, + } + toolMetadataByName := map[string]einocallbacks.ToolMetadata{ + toolx.GraphPrepareTicketDraft.Name: {ToolCode: toolx.GraphPrepareTicketDraft.Code, ToolName: toolx.GraphPrepareTicketDraft.Name}, + "mcp_server_refund": {ToolCode: "mcp/refund", ToolName: "mcp_server_refund"}, + } + got := extractToolCodesFromInfos(infos, toolMetadataByName) + if len(got) != 3 { + t.Fatalf("unexpected tool codes: %#v", got) + } + if got[0] != toolx.BuiltinSkill.Code || got[1] != toolx.GraphPrepareTicketDraft.Code || got[2] != "mcp/refund" { + t.Fatalf("unexpected tool codes order: %#v", got) + } +}