package services import ( "context" "encoding/json" "fmt" "path/filepath" "runtime" "strings" "sync" "time" "server/models" mcpclient "github.com/mark3labs/mcp-go/client" "github.com/mark3labs/mcp-go/client/transport" "github.com/mark3labs/mcp-go/mcp" ) // McpSession 一次已建立的 MCP 连接会话 type McpSession struct { Server models.BackendMcpServer Client *mcpclient.Client Tools []models.McpToolInfo Finger string // 配置指纹,配置变更时自动重连 } // McpManager MCP 客户端管理器(全局单例,带连接缓存) type McpManager struct { mu sync.Mutex sessions map[uint64]*McpSession } // NewMcpManager 创建管理器 func NewMcpManager() *McpManager { return &McpManager{sessions: make(map[uint64]*McpSession)} } // McpClientManager 全局 MCP 客户端管理器 var McpClientManager = NewMcpManager() const ( connectTimeout = 25 * time.Second callTimeout = 90 * time.Second ) // serverFinger 计算服务器配置指纹(用于配置变更自动重连) func serverFinger(s *models.BackendMcpServer) string { return fmt.Sprintf("%s|%s|%s|%s|%s|%s|%s", s.Transport, s.Command, s.Args, s.Env, s.URL, s.Headers, s.Name) } // parseStringArray 解析 JSON 数组字符串为 []string func parseStringArray(s string) []string { s = strings.TrimSpace(s) if s == "" { return nil } var arr []string if err := json.Unmarshal([]byte(s), &arr); err == nil { return arr } // 兼容逗号分隔 parts := strings.Split(s, ",") out := make([]string, 0, len(parts)) for _, p := range parts { if v := strings.TrimSpace(p); v != "" { out = append(out, v) } } return out } // parseStringMap 解析 JSON 对象字符串为 map func parseStringMap(s string) map[string]string { s = strings.TrimSpace(s) if s == "" { return nil } var m map[string]string if err := json.Unmarshal([]byte(s), &m); err == nil { return m } var raw map[string]interface{} if err := json.Unmarshal([]byte(s), &raw); err == nil { out := make(map[string]string, len(raw)) for k, v := range raw { out[k] = fmt.Sprintf("%v", v) } return out } return nil } // envToSlice 将环境变量 map 转为 "K=V" 切片 func envToSlice(env map[string]string) []string { if len(env) == 0 { return nil } out := make([]string, 0, len(env)) for k, v := range env { out = append(out, k+"="+v) } return out } // resolveStdioCommand Windows 下 .cmd/.bat/npx 需要 cmd /c 包装 func resolveStdioCommand(cmd string, args []string) (string, []string) { if runtime.GOOS != "windows" { return cmd, args } lower := strings.ToLower(strings.TrimSpace(cmd)) base := filepath.Base(lower) if base == "npx" || base == "npm" || base == "npx.cmd" || base == "npm.cmd" || base == "uvx" || base == "uvx.exe" || strings.HasSuffix(lower, ".cmd") || strings.HasSuffix(lower, ".bat") { all := append([]string{cmd}, args...) return "cmd", append([]string{"/c"}, all...) } return cmd, args } // buildClient 按传输类型创建 MCP 客户端(不连接、不初始化) func buildClient(s *models.BackendMcpServer) (*mcpclient.Client, error) { switch s.Transport { case "stdio": cmd, args := resolveStdioCommand(strings.TrimSpace(s.Command), parseStringArray(s.Args)) if cmd == "" { return nil, fmt.Errorf("stdio 传输必须配置 command") } return mcpclient.NewStdioMCPClient(cmd, envToSlice(parseStringMap(s.Env)), args...) case "sse": if strings.TrimSpace(s.URL) == "" { return nil, fmt.Errorf("sse 传输必须配置 url") } return mcpclient.NewSSEMCPClient(strings.TrimSpace(s.URL), transport.WithHeaders(parseStringMap(s.Headers))) case "http": if strings.TrimSpace(s.URL) == "" { return nil, fmt.Errorf("http 传输必须配置 url") } return mcpclient.NewStreamableHttpClient(strings.TrimSpace(s.URL), transport.WithHTTPHeaders(parseStringMap(s.Headers))) default: return nil, fmt.Errorf("不支持的传输类型: %s", s.Transport) } } // connectAndList 建立连接并列出工具 func connectAndList(ctx context.Context, s *models.BackendMcpServer) (*mcpclient.Client, []models.McpToolInfo, error) { cl, err := buildClient(s) if err != nil { return nil, nil, err } if s.Transport != "stdio" { // stdio 的传输在构造函数中已启动,其余需手动 Start if err := cl.Start(ctx); err != nil { _ = cl.Close() return nil, nil, fmt.Errorf("启动连接失败: %w", err) } } initReq := mcp.InitializeRequest{} initReq.Params.ProtocolVersion = mcp.LATEST_PROTOCOL_VERSION initReq.Params.ClientInfo = mcp.Implementation{ Name: "xiaozhi-ai-backend", Version: "1.0.0", } if _, err := cl.Initialize(ctx, initReq); err != nil { _ = cl.Close() return nil, nil, fmt.Errorf("MCP 握手失败: %w", err) } toolsResult, err := cl.ListTools(ctx, mcp.ListToolsRequest{}) if err != nil { _ = cl.Close() return nil, nil, fmt.Errorf("获取工具列表失败: %w", err) } tools := make([]models.McpToolInfo, 0, len(toolsResult.Tools)) for _, t := range toolsResult.Tools { info := models.McpToolInfo{ Name: t.Name, Description: t.Description, } // 优先使用 RawInputSchema(完整 JSON Schema) if len(t.RawInputSchema) > 0 { var schema interface{} if err := json.Unmarshal(t.RawInputSchema, &schema); err == nil { info.InputSchema = schema } } else if t.InputSchema.Type != "" || len(t.InputSchema.Properties) > 0 { info.InputSchema = t.InputSchema } tools = append(tools, info) } return cl, tools, nil } // TestConnection 测试连接:新建连接 → 列出工具 → 关闭(不进入缓存) func (m *McpManager) TestConnection(s *models.BackendMcpServer) ([]models.McpToolInfo, error) { ctx, cancel := context.WithTimeout(context.Background(), connectTimeout) defer cancel() cl, tools, err := connectAndList(ctx, s) if err != nil { return nil, err } _ = cl.Close() return tools, nil } // EnsureConnected 获取(或建立)连接,返回该服务器已发现的工具 func (m *McpManager) EnsureConnected(s *models.BackendMcpServer) ([]models.McpToolInfo, error) { finger := serverFinger(s) m.mu.Lock() if sess, ok := m.sessions[s.ID]; ok && sess.Finger == finger { tools := sess.Tools m.mu.Unlock() return tools, nil } m.mu.Unlock() // 新建连接(放在锁外,避免长时间占用锁) ctx, cancel := context.WithTimeout(context.Background(), connectTimeout) defer cancel() cl, tools, err := connectAndList(ctx, s) if err != nil { return nil, err } m.mu.Lock() defer m.mu.Unlock() // 关闭旧连接 if old, ok := m.sessions[s.ID]; ok && old.Client != nil { _ = old.Client.Close() } m.sessions[s.ID] = &McpSession{ Server: *s, Client: cl, Tools: tools, Finger: finger, } return tools, nil } // ListTools 返回缓存中的工具列表(未连接返回 nil) func (m *McpManager) ListTools(serverID uint64) []models.McpToolInfo { m.mu.Lock() defer m.mu.Unlock() if sess, ok := m.sessions[serverID]; ok { return sess.Tools } return nil } // CallTool 调用 MCP 工具,返回文本结果 func (m *McpManager) CallTool(serverID uint64, name string, args map[string]interface{}) (string, bool, error) { m.mu.Lock() sess, ok := m.sessions[serverID] m.mu.Unlock() if !ok || sess.Client == nil { return "", false, fmt.Errorf("MCP 服务未连接") } ctx, cancel := context.WithTimeout(context.Background(), callTimeout) defer cancel() req := mcp.CallToolRequest{Params: mcp.CallToolParams{ Name: name, Arguments: args, }} result, err := sess.Client.CallTool(ctx, req) if err != nil { return "", false, err } return mcpResultToText(result), result.IsError, nil } // mcpResultToText 将 MCP CallToolResult 内容转为文本 func mcpResultToText(result *mcp.CallToolResult) string { if result == nil { return "" } var sb strings.Builder for _, c := range result.Content { switch v := c.(type) { case mcp.TextContent: sb.WriteString(v.Text) case mcp.ImageContent: sb.WriteString(fmt.Sprintf("[图片: %s, %d 字节]", v.MIMEType, len(v.Data))) case mcp.AudioContent: sb.WriteString(fmt.Sprintf("[音频: %s, %d 字节]", v.MIMEType, len(v.Data))) default: if b, err := json.Marshal(c); err == nil { sb.Write(b) } } } return sb.String() } // Disconnect 断开并移除指定服务器的连接 func (m *McpManager) Disconnect(serverID uint64) { m.mu.Lock() defer m.mu.Unlock() if sess, ok := m.sessions[serverID]; ok { if sess.Client != nil { _ = sess.Client.Close() } delete(m.sessions, serverID) } } // CloseAll 关闭所有连接(服务退出时调用) func (m *McpManager) CloseAll() { m.mu.Lock() defer m.mu.Unlock() for id, sess := range m.sessions { if sess.Client != nil { _ = sess.Client.Close() } delete(m.sessions, id) } }