Files
2026-09-20 00:19:08 +08:00

326 lines
8.6 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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)
}
}