diff --git a/internal/einomcp/mcp_tools.go b/internal/einomcp/mcp_tools.go index edff81b4a..d63d0ca62 100644 --- a/internal/einomcp/mcp_tools.go +++ b/internal/einomcp/mcp_tools.go @@ -45,14 +45,14 @@ func ToolsFromDefinitions( return nil, fmt.Errorf("tool %q: %w", d.Function.Name, err) } out = append(out, &mcpBridgeTool{ - info: info, - name: d.Function.Name, - agent: ag, - holder: holder, - record: rec, - chunk: toolOutputChunk, - invokeNotify: invokeNotify, - einoAgentName: strings.TrimSpace(einoAgentName), + info: info, + name: d.Function.Name, + agent: ag, + holder: holder, + record: rec, + chunk: toolOutputChunk, + invokeNotify: invokeNotify, + einoAgentName: strings.TrimSpace(einoAgentName), }) } return out, nil @@ -77,12 +77,29 @@ func toolInfoFromDefinition(d agent.Tool) (*schema.ToolInfo, error) { // 空参数对象 } return &schema.ToolInfo{ - Name: fn.Name, + Name: sanitizeOpenAIToolName(fn.Name), Desc: fn.Description, ParamsOneOf: schema.NewParamsOneOfByJSONSchema(&js), }, nil } +// sanitizeOpenAIToolName converts MCP names to OpenAI's ^[a-zA-Z0-9_-]+$ format. +// The original name remains on mcpBridgeTool and is still used for MCP routing. +func sanitizeOpenAIToolName(name string) string { + name = strings.ReplaceAll(name, "::", "__") + var b strings.Builder + b.Grow(len(name)) + for _, r := range name { + if (r >= 'a' && r <= 'z') || (r >= 'A' && r <= 'Z') || + (r >= '0' && r <= '9') || r == '_' || r == '-' { + b.WriteRune(r) + } else { + b.WriteByte('_') + } + } + return b.String() +} + type mcpBridgeTool struct { info *schema.ToolInfo name string diff --git a/internal/einomcp/mcp_tools_test.go b/internal/einomcp/mcp_tools_test.go index 078c8c04e..3487e611a 100644 --- a/internal/einomcp/mcp_tools_test.go +++ b/internal/einomcp/mcp_tools_test.go @@ -3,8 +3,40 @@ package einomcp import ( "strings" "testing" + + "cyberstrike-ai/internal/agent" ) +func TestToolInfoFromDefinitionSanitizesOpenAIToolName(t *testing.T) { + tests := []struct { + name string + want string + }{ + {name: "fs.read", want: "fs_read"}, + {name: "nezha::server.exec", want: "nezha__server_exec"}, + {name: "already-valid_1", want: "already-valid_1"}, + {name: "space/slash", want: "space_slash"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + info, err := toolInfoFromDefinition(agent.Tool{ + Type: "function", + Function: agent.FunctionDefinition{ + Name: tt.name, + Parameters: map[string]interface{}{"type": "object"}, + }, + }) + if err != nil { + t.Fatalf("toolInfoFromDefinition() error = %v", err) + } + if info.Name != tt.want { + t.Fatalf("toolInfoFromDefinition() name = %q, want %q", info.Name, tt.want) + } + }) + } +} + func TestUnknownToolReminderText(t *testing.T) { s := unknownToolReminderText("bad_tool") if !strings.Contains(s, "bad_tool") {