增加供应商智能添加功能
This commit is contained in:
@@ -0,0 +1,325 @@
|
||||
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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user