326 lines
8.6 KiB
Go
326 lines
8.6 KiB
Go
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)
|
||
}
|
||
}
|