mirror of
https://github.com/Ed1s0nZ/CyberStrikeAI.git
synced 2026-08-13 22:50:24 +02:00
Add files via upload
This commit is contained in:
@@ -0,0 +1,474 @@
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"os/exec"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
// ExternalMCPClient 外部MCP客户端接口
|
||||
type ExternalMCPClient interface {
|
||||
// Initialize 初始化连接
|
||||
Initialize(ctx context.Context) error
|
||||
// ListTools 列出工具
|
||||
ListTools(ctx context.Context) ([]Tool, error)
|
||||
// CallTool 调用工具
|
||||
CallTool(ctx context.Context, name string, args map[string]interface{}) (*ToolResult, error)
|
||||
// Close 关闭连接
|
||||
Close() error
|
||||
// IsConnected 检查是否已连接
|
||||
IsConnected() bool
|
||||
// GetStatus 获取状态
|
||||
GetStatus() string
|
||||
}
|
||||
|
||||
// HTTPMCPClient HTTP模式的MCP客户端
|
||||
type HTTPMCPClient struct {
|
||||
url string
|
||||
timeout time.Duration
|
||||
client *http.Client
|
||||
logger *zap.Logger
|
||||
mu sync.RWMutex
|
||||
status string // "disconnected", "connecting", "connected", "error"
|
||||
}
|
||||
|
||||
// NewHTTPMCPClient 创建HTTP模式的MCP客户端
|
||||
func NewHTTPMCPClient(url string, timeout time.Duration, logger *zap.Logger) *HTTPMCPClient {
|
||||
if timeout <= 0 {
|
||||
timeout = 30 * time.Second
|
||||
}
|
||||
return &HTTPMCPClient{
|
||||
url: url,
|
||||
timeout: timeout,
|
||||
client: &http.Client{
|
||||
Timeout: timeout,
|
||||
},
|
||||
logger: logger,
|
||||
status: "disconnected",
|
||||
}
|
||||
}
|
||||
|
||||
func (c *HTTPMCPClient) setStatus(status string) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
c.status = status
|
||||
}
|
||||
|
||||
func (c *HTTPMCPClient) GetStatus() string {
|
||||
c.mu.RLock()
|
||||
defer c.mu.RUnlock()
|
||||
return c.status
|
||||
}
|
||||
|
||||
func (c *HTTPMCPClient) IsConnected() bool {
|
||||
return c.GetStatus() == "connected"
|
||||
}
|
||||
|
||||
func (c *HTTPMCPClient) Initialize(ctx context.Context) error {
|
||||
c.setStatus("connecting")
|
||||
|
||||
req := Message{
|
||||
ID: MessageID{value: "1"},
|
||||
Method: "initialize",
|
||||
Version: "2.0",
|
||||
}
|
||||
|
||||
params := InitializeRequest{
|
||||
ProtocolVersion: ProtocolVersion,
|
||||
Capabilities: make(map[string]interface{}),
|
||||
ClientInfo: ClientInfo{
|
||||
Name: "CyberStrikeAI",
|
||||
Version: "1.0.0",
|
||||
},
|
||||
}
|
||||
|
||||
paramsJSON, _ := json.Marshal(params)
|
||||
req.Params = paramsJSON
|
||||
|
||||
_, err := c.sendRequest(ctx, &req)
|
||||
if err != nil {
|
||||
c.setStatus("error")
|
||||
return fmt.Errorf("初始化失败: %w", err)
|
||||
}
|
||||
|
||||
c.setStatus("connected")
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *HTTPMCPClient) ListTools(ctx context.Context) ([]Tool, error) {
|
||||
req := Message{
|
||||
ID: MessageID{value: uuid.New().String()},
|
||||
Method: "tools/list",
|
||||
Version: "2.0",
|
||||
}
|
||||
|
||||
req.Params = json.RawMessage("{}")
|
||||
|
||||
resp, err := c.sendRequest(ctx, &req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("获取工具列表失败: %w", err)
|
||||
}
|
||||
|
||||
var listResp ListToolsResponse
|
||||
if err := json.Unmarshal(resp.Result, &listResp); err != nil {
|
||||
return nil, fmt.Errorf("解析工具列表失败: %w", err)
|
||||
}
|
||||
|
||||
return listResp.Tools, nil
|
||||
}
|
||||
|
||||
func (c *HTTPMCPClient) CallTool(ctx context.Context, name string, args map[string]interface{}) (*ToolResult, error) {
|
||||
req := Message{
|
||||
ID: MessageID{value: uuid.New().String()},
|
||||
Method: "tools/call",
|
||||
Version: "2.0",
|
||||
}
|
||||
|
||||
callReq := CallToolRequest{
|
||||
Name: name,
|
||||
Arguments: args,
|
||||
}
|
||||
|
||||
paramsJSON, _ := json.Marshal(callReq)
|
||||
req.Params = paramsJSON
|
||||
|
||||
resp, err := c.sendRequest(ctx, &req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("调用工具失败: %w", err)
|
||||
}
|
||||
|
||||
var callResp CallToolResponse
|
||||
if err := json.Unmarshal(resp.Result, &callResp); err != nil {
|
||||
return nil, fmt.Errorf("解析工具调用结果失败: %w", err)
|
||||
}
|
||||
|
||||
return &ToolResult{
|
||||
Content: callResp.Content,
|
||||
IsError: callResp.IsError,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (c *HTTPMCPClient) sendRequest(ctx context.Context, msg *Message) (*Message, error) {
|
||||
body, err := json.Marshal(msg)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("序列化请求失败: %w", err)
|
||||
}
|
||||
|
||||
httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, c.url, bytes.NewReader(body))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("创建HTTP请求失败: %w", err)
|
||||
}
|
||||
|
||||
httpReq.Header.Set("Content-Type", "application/json")
|
||||
|
||||
resp, err := c.client.Do(httpReq)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("HTTP请求失败: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
bodyBytes, _ := io.ReadAll(resp.Body)
|
||||
return nil, fmt.Errorf("HTTP错误 %d: %s", resp.StatusCode, string(bodyBytes))
|
||||
}
|
||||
|
||||
var mcpResp Message
|
||||
if err := json.NewDecoder(resp.Body).Decode(&mcpResp); err != nil {
|
||||
return nil, fmt.Errorf("解析响应失败: %w", err)
|
||||
}
|
||||
|
||||
if mcpResp.Error != nil {
|
||||
return nil, fmt.Errorf("MCP错误: %s (code: %d)", mcpResp.Error.Message, mcpResp.Error.Code)
|
||||
}
|
||||
|
||||
return &mcpResp, nil
|
||||
}
|
||||
|
||||
func (c *HTTPMCPClient) Close() error {
|
||||
c.setStatus("disconnected")
|
||||
return nil
|
||||
}
|
||||
|
||||
// StdioMCPClient stdio模式的MCP客户端
|
||||
type StdioMCPClient struct {
|
||||
command string
|
||||
args []string
|
||||
timeout time.Duration
|
||||
cmd *exec.Cmd
|
||||
stdin io.WriteCloser
|
||||
stdout io.ReadCloser
|
||||
decoder *json.Decoder
|
||||
encoder *json.Encoder
|
||||
logger *zap.Logger
|
||||
mu sync.RWMutex
|
||||
status string
|
||||
requestID int64
|
||||
responses map[string]chan *Message
|
||||
responsesMu sync.Mutex
|
||||
ctx context.Context
|
||||
cancel context.CancelFunc
|
||||
}
|
||||
|
||||
// NewStdioMCPClient 创建stdio模式的MCP客户端
|
||||
func NewStdioMCPClient(command string, args []string, timeout time.Duration, logger *zap.Logger) *StdioMCPClient {
|
||||
if timeout <= 0 {
|
||||
timeout = 30 * time.Second
|
||||
}
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
return &StdioMCPClient{
|
||||
command: command,
|
||||
args: args,
|
||||
timeout: timeout,
|
||||
logger: logger,
|
||||
status: "disconnected",
|
||||
responses: make(map[string]chan *Message),
|
||||
ctx: ctx,
|
||||
cancel: cancel,
|
||||
}
|
||||
}
|
||||
|
||||
func (c *StdioMCPClient) setStatus(status string) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
c.status = status
|
||||
}
|
||||
|
||||
func (c *StdioMCPClient) GetStatus() string {
|
||||
c.mu.RLock()
|
||||
defer c.mu.RUnlock()
|
||||
return c.status
|
||||
}
|
||||
|
||||
func (c *StdioMCPClient) IsConnected() bool {
|
||||
return c.GetStatus() == "connected"
|
||||
}
|
||||
|
||||
func (c *StdioMCPClient) Initialize(ctx context.Context) error {
|
||||
c.setStatus("connecting")
|
||||
|
||||
if err := c.startProcess(); err != nil {
|
||||
c.setStatus("error")
|
||||
return fmt.Errorf("启动进程失败: %w", err)
|
||||
}
|
||||
|
||||
// 启动响应读取goroutine
|
||||
go c.readResponses()
|
||||
|
||||
// 发送初始化请求
|
||||
req := Message{
|
||||
ID: MessageID{value: "1"},
|
||||
Method: "initialize",
|
||||
Version: "2.0",
|
||||
}
|
||||
|
||||
params := InitializeRequest{
|
||||
ProtocolVersion: ProtocolVersion,
|
||||
Capabilities: make(map[string]interface{}),
|
||||
ClientInfo: ClientInfo{
|
||||
Name: "CyberStrikeAI",
|
||||
Version: "1.0.0",
|
||||
},
|
||||
}
|
||||
|
||||
paramsJSON, _ := json.Marshal(params)
|
||||
req.Params = paramsJSON
|
||||
|
||||
_, err := c.sendRequest(ctx, &req)
|
||||
if err != nil {
|
||||
c.setStatus("error")
|
||||
c.Close()
|
||||
return fmt.Errorf("初始化失败: %w", err)
|
||||
}
|
||||
|
||||
c.setStatus("connected")
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *StdioMCPClient) startProcess() error {
|
||||
cmd := exec.CommandContext(c.ctx, c.command, c.args...)
|
||||
|
||||
stdin, err := cmd.StdinPipe()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
stdout, err := cmd.StdoutPipe()
|
||||
if err != nil {
|
||||
stdin.Close()
|
||||
return err
|
||||
}
|
||||
|
||||
if err := cmd.Start(); err != nil {
|
||||
stdin.Close()
|
||||
stdout.Close()
|
||||
return err
|
||||
}
|
||||
|
||||
c.cmd = cmd
|
||||
c.stdin = stdin
|
||||
c.stdout = stdout
|
||||
c.decoder = json.NewDecoder(stdout)
|
||||
c.encoder = json.NewEncoder(stdin)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *StdioMCPClient) readResponses() {
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
c.logger.Error("读取响应时发生panic", zap.Any("error", r))
|
||||
}
|
||||
}()
|
||||
|
||||
for {
|
||||
var msg Message
|
||||
if err := c.decoder.Decode(&msg); err != nil {
|
||||
if err == io.EOF {
|
||||
c.setStatus("disconnected")
|
||||
break
|
||||
}
|
||||
c.logger.Error("读取响应失败", zap.Error(err))
|
||||
break
|
||||
}
|
||||
|
||||
// 处理响应
|
||||
id := msg.ID.String()
|
||||
c.responsesMu.Lock()
|
||||
if ch, ok := c.responses[id]; ok {
|
||||
select {
|
||||
case ch <- &msg:
|
||||
default:
|
||||
}
|
||||
delete(c.responses, id)
|
||||
}
|
||||
c.responsesMu.Unlock()
|
||||
}
|
||||
}
|
||||
|
||||
func (c *StdioMCPClient) sendRequest(ctx context.Context, msg *Message) (*Message, error) {
|
||||
if c.encoder == nil {
|
||||
return nil, fmt.Errorf("进程未启动")
|
||||
}
|
||||
|
||||
id := msg.ID.String()
|
||||
if id == "" {
|
||||
c.mu.Lock()
|
||||
c.requestID++
|
||||
id = fmt.Sprintf("%d", c.requestID)
|
||||
msg.ID = MessageID{value: id}
|
||||
c.mu.Unlock()
|
||||
}
|
||||
|
||||
// 创建响应通道
|
||||
responseCh := make(chan *Message, 1)
|
||||
c.responsesMu.Lock()
|
||||
c.responses[id] = responseCh
|
||||
c.responsesMu.Unlock()
|
||||
|
||||
// 发送请求
|
||||
if err := c.encoder.Encode(msg); err != nil {
|
||||
c.responsesMu.Lock()
|
||||
delete(c.responses, id)
|
||||
c.responsesMu.Unlock()
|
||||
return nil, fmt.Errorf("发送请求失败: %w", err)
|
||||
}
|
||||
|
||||
// 等待响应
|
||||
select {
|
||||
case resp := <-responseCh:
|
||||
if resp.Error != nil {
|
||||
return nil, fmt.Errorf("MCP错误: %s (code: %d)", resp.Error.Message, resp.Error.Code)
|
||||
}
|
||||
return resp, nil
|
||||
case <-ctx.Done():
|
||||
c.responsesMu.Lock()
|
||||
delete(c.responses, id)
|
||||
c.responsesMu.Unlock()
|
||||
return nil, ctx.Err()
|
||||
case <-time.After(c.timeout):
|
||||
c.responsesMu.Lock()
|
||||
delete(c.responses, id)
|
||||
c.responsesMu.Unlock()
|
||||
return nil, fmt.Errorf("请求超时")
|
||||
}
|
||||
}
|
||||
|
||||
func (c *StdioMCPClient) ListTools(ctx context.Context) ([]Tool, error) {
|
||||
req := Message{
|
||||
ID: MessageID{value: uuid.New().String()},
|
||||
Method: "tools/list",
|
||||
Version: "2.0",
|
||||
}
|
||||
|
||||
req.Params = json.RawMessage("{}")
|
||||
|
||||
resp, err := c.sendRequest(ctx, &req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("获取工具列表失败: %w", err)
|
||||
}
|
||||
|
||||
var listResp ListToolsResponse
|
||||
if err := json.Unmarshal(resp.Result, &listResp); err != nil {
|
||||
return nil, fmt.Errorf("解析工具列表失败: %w", err)
|
||||
}
|
||||
|
||||
return listResp.Tools, nil
|
||||
}
|
||||
|
||||
func (c *StdioMCPClient) CallTool(ctx context.Context, name string, args map[string]interface{}) (*ToolResult, error) {
|
||||
req := Message{
|
||||
ID: MessageID{value: uuid.New().String()},
|
||||
Method: "tools/call",
|
||||
Version: "2.0",
|
||||
}
|
||||
|
||||
callReq := CallToolRequest{
|
||||
Name: name,
|
||||
Arguments: args,
|
||||
}
|
||||
|
||||
paramsJSON, _ := json.Marshal(callReq)
|
||||
req.Params = paramsJSON
|
||||
|
||||
resp, err := c.sendRequest(ctx, &req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("调用工具失败: %w", err)
|
||||
}
|
||||
|
||||
var callResp CallToolResponse
|
||||
if err := json.Unmarshal(resp.Result, &callResp); err != nil {
|
||||
return nil, fmt.Errorf("解析工具调用结果失败: %w", err)
|
||||
}
|
||||
|
||||
return &ToolResult{
|
||||
Content: callResp.Content,
|
||||
IsError: callResp.IsError,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (c *StdioMCPClient) Close() error {
|
||||
c.cancel()
|
||||
|
||||
if c.stdin != nil {
|
||||
c.stdin.Close()
|
||||
}
|
||||
if c.stdout != nil {
|
||||
c.stdout.Close()
|
||||
}
|
||||
if c.cmd != nil {
|
||||
c.cmd.Process.Kill()
|
||||
c.cmd.Wait()
|
||||
}
|
||||
|
||||
c.setStatus("disconnected")
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,660 @@
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"cyberstrike-ai/internal/config"
|
||||
"github.com/google/uuid"
|
||||
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
// ExternalMCPManager 外部MCP管理器
|
||||
type ExternalMCPManager struct {
|
||||
clients map[string]ExternalMCPClient
|
||||
configs map[string]config.ExternalMCPServerConfig
|
||||
logger *zap.Logger
|
||||
storage MonitorStorage // 可选的持久化存储
|
||||
executions map[string]*ToolExecution // 执行记录
|
||||
stats map[string]*ToolStats // 工具统计信息
|
||||
mu sync.RWMutex
|
||||
}
|
||||
|
||||
// NewExternalMCPManager 创建外部MCP管理器
|
||||
func NewExternalMCPManager(logger *zap.Logger) *ExternalMCPManager {
|
||||
return NewExternalMCPManagerWithStorage(logger, nil)
|
||||
}
|
||||
|
||||
// NewExternalMCPManagerWithStorage 创建外部MCP管理器(带持久化存储)
|
||||
func NewExternalMCPManagerWithStorage(logger *zap.Logger, storage MonitorStorage) *ExternalMCPManager {
|
||||
return &ExternalMCPManager{
|
||||
clients: make(map[string]ExternalMCPClient),
|
||||
configs: make(map[string]config.ExternalMCPServerConfig),
|
||||
logger: logger,
|
||||
storage: storage,
|
||||
executions: make(map[string]*ToolExecution),
|
||||
stats: make(map[string]*ToolStats),
|
||||
}
|
||||
}
|
||||
|
||||
// LoadConfigs 加载配置
|
||||
func (m *ExternalMCPManager) LoadConfigs(cfg *config.ExternalMCPConfig) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
if cfg == nil || cfg.Servers == nil {
|
||||
return
|
||||
}
|
||||
|
||||
m.configs = make(map[string]config.ExternalMCPServerConfig)
|
||||
for name, serverCfg := range cfg.Servers {
|
||||
m.configs[name] = serverCfg
|
||||
}
|
||||
}
|
||||
|
||||
// GetConfigs 获取所有配置
|
||||
func (m *ExternalMCPManager) GetConfigs() map[string]config.ExternalMCPServerConfig {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
|
||||
result := make(map[string]config.ExternalMCPServerConfig)
|
||||
for k, v := range m.configs {
|
||||
result[k] = v
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// AddOrUpdateConfig 添加或更新配置
|
||||
func (m *ExternalMCPManager) AddOrUpdateConfig(name string, serverCfg config.ExternalMCPServerConfig) error {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
// 如果已存在客户端,先关闭
|
||||
if client, exists := m.clients[name]; exists {
|
||||
client.Close()
|
||||
delete(m.clients, name)
|
||||
}
|
||||
|
||||
m.configs[name] = serverCfg
|
||||
|
||||
// 如果启用,自动连接
|
||||
if m.isEnabled(serverCfg) {
|
||||
go m.connectClient(name, serverCfg)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// RemoveConfig 移除配置
|
||||
func (m *ExternalMCPManager) RemoveConfig(name string) error {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
// 关闭客户端
|
||||
if client, exists := m.clients[name]; exists {
|
||||
client.Close()
|
||||
delete(m.clients, name)
|
||||
}
|
||||
|
||||
delete(m.configs, name)
|
||||
return nil
|
||||
}
|
||||
|
||||
// StartClient 启动客户端
|
||||
func (m *ExternalMCPManager) StartClient(name string) error {
|
||||
m.mu.Lock()
|
||||
serverCfg, exists := m.configs[name]
|
||||
m.mu.Unlock()
|
||||
|
||||
if !exists {
|
||||
return fmt.Errorf("配置不存在: %s", name)
|
||||
}
|
||||
|
||||
// 检查是否已经有连接的客户端
|
||||
m.mu.RLock()
|
||||
_, hasClient := m.clients[name]
|
||||
m.mu.RUnlock()
|
||||
|
||||
if hasClient {
|
||||
// 检查客户端是否已连接
|
||||
if client, ok := m.GetClient(name); ok && client.IsConnected() {
|
||||
return fmt.Errorf("客户端已连接")
|
||||
}
|
||||
// 如果有客户端但未连接,先关闭
|
||||
if client, ok := m.GetClient(name); ok {
|
||||
client.Close()
|
||||
m.mu.Lock()
|
||||
delete(m.clients, name)
|
||||
m.mu.Unlock()
|
||||
}
|
||||
}
|
||||
|
||||
// 更新配置为启用
|
||||
m.mu.Lock()
|
||||
serverCfg.ExternalMCPEnable = true
|
||||
m.configs[name] = serverCfg
|
||||
m.mu.Unlock()
|
||||
|
||||
// 连接客户端
|
||||
return m.connectClient(name, serverCfg)
|
||||
}
|
||||
|
||||
// StopClient 停止客户端
|
||||
func (m *ExternalMCPManager) StopClient(name string) error {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
serverCfg, exists := m.configs[name]
|
||||
if !exists {
|
||||
return fmt.Errorf("配置不存在: %s", name)
|
||||
}
|
||||
|
||||
// 关闭客户端
|
||||
if client, exists := m.clients[name]; exists {
|
||||
client.Close()
|
||||
delete(m.clients, name)
|
||||
}
|
||||
|
||||
// 更新配置为禁用
|
||||
serverCfg.ExternalMCPEnable = false
|
||||
m.configs[name] = serverCfg
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetClient 获取客户端
|
||||
func (m *ExternalMCPManager) GetClient(name string) (ExternalMCPClient, bool) {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
|
||||
client, exists := m.clients[name]
|
||||
return client, exists
|
||||
}
|
||||
|
||||
// GetAllTools 获取所有外部MCP的工具
|
||||
func (m *ExternalMCPManager) GetAllTools(ctx context.Context) ([]Tool, error) {
|
||||
m.mu.RLock()
|
||||
clients := make(map[string]ExternalMCPClient)
|
||||
for k, v := range m.clients {
|
||||
clients[k] = v
|
||||
}
|
||||
m.mu.RUnlock()
|
||||
|
||||
var allTools []Tool
|
||||
for name, client := range clients {
|
||||
if !client.IsConnected() {
|
||||
continue
|
||||
}
|
||||
|
||||
tools, err := client.ListTools(ctx)
|
||||
if err != nil {
|
||||
m.logger.Warn("获取外部MCP工具列表失败",
|
||||
zap.String("name", name),
|
||||
zap.Error(err),
|
||||
)
|
||||
continue
|
||||
}
|
||||
|
||||
// 为工具添加前缀,避免冲突
|
||||
for _, tool := range tools {
|
||||
tool.Name = fmt.Sprintf("%s::%s", name, tool.Name)
|
||||
allTools = append(allTools, tool)
|
||||
}
|
||||
}
|
||||
|
||||
return allTools, nil
|
||||
}
|
||||
|
||||
// CallTool 调用外部MCP工具(返回执行ID)
|
||||
func (m *ExternalMCPManager) CallTool(ctx context.Context, toolName string, args map[string]interface{}) (*ToolResult, string, error) {
|
||||
// 解析工具名称:name::toolName
|
||||
var mcpName, actualToolName string
|
||||
if idx := findSubstring(toolName, "::"); idx > 0 {
|
||||
mcpName = toolName[:idx]
|
||||
actualToolName = toolName[idx+2:]
|
||||
} else {
|
||||
return nil, "", fmt.Errorf("无效的工具名称格式: %s", toolName)
|
||||
}
|
||||
|
||||
client, exists := m.GetClient(mcpName)
|
||||
if !exists {
|
||||
return nil, "", fmt.Errorf("外部MCP客户端不存在: %s", mcpName)
|
||||
}
|
||||
|
||||
if !client.IsConnected() {
|
||||
return nil, "", fmt.Errorf("外部MCP客户端未连接: %s", mcpName)
|
||||
}
|
||||
|
||||
// 创建执行记录
|
||||
executionID := uuid.New().String()
|
||||
execution := &ToolExecution{
|
||||
ID: executionID,
|
||||
ToolName: toolName, // 使用完整工具名称(包含MCP名称)
|
||||
Arguments: args,
|
||||
Status: "running",
|
||||
StartTime: time.Now(),
|
||||
}
|
||||
|
||||
m.mu.Lock()
|
||||
m.executions[executionID] = execution
|
||||
// 如果内存中的执行记录超过限制,清理最旧的记录
|
||||
m.cleanupOldExecutions()
|
||||
m.mu.Unlock()
|
||||
|
||||
if m.storage != nil {
|
||||
if err := m.storage.SaveToolExecution(execution); err != nil {
|
||||
m.logger.Warn("保存执行记录到数据库失败", zap.Error(err))
|
||||
}
|
||||
}
|
||||
|
||||
// 调用工具
|
||||
result, err := client.CallTool(ctx, actualToolName, args)
|
||||
|
||||
// 更新执行记录
|
||||
m.mu.Lock()
|
||||
now := time.Now()
|
||||
execution.EndTime = &now
|
||||
execution.Duration = now.Sub(execution.StartTime)
|
||||
|
||||
if err != nil {
|
||||
execution.Status = "failed"
|
||||
execution.Error = err.Error()
|
||||
} else if result != nil && result.IsError {
|
||||
execution.Status = "failed"
|
||||
if len(result.Content) > 0 {
|
||||
execution.Error = result.Content[0].Text
|
||||
} else {
|
||||
execution.Error = "工具执行返回错误结果"
|
||||
}
|
||||
execution.Result = result
|
||||
} else {
|
||||
execution.Status = "completed"
|
||||
if result == nil {
|
||||
result = &ToolResult{
|
||||
Content: []Content{
|
||||
{Type: "text", Text: "工具执行完成,但未返回结果"},
|
||||
},
|
||||
}
|
||||
}
|
||||
execution.Result = result
|
||||
}
|
||||
m.mu.Unlock()
|
||||
|
||||
if m.storage != nil {
|
||||
if err := m.storage.SaveToolExecution(execution); err != nil {
|
||||
m.logger.Warn("保存执行记录到数据库失败", zap.Error(err))
|
||||
}
|
||||
}
|
||||
|
||||
// 更新统计信息
|
||||
failed := err != nil || (result != nil && result.IsError)
|
||||
m.updateStats(toolName, failed)
|
||||
|
||||
// 如果使用存储,从内存中删除(已持久化)
|
||||
if m.storage != nil {
|
||||
m.mu.Lock()
|
||||
delete(m.executions, executionID)
|
||||
m.mu.Unlock()
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
return nil, executionID, err
|
||||
}
|
||||
|
||||
return result, executionID, nil
|
||||
}
|
||||
|
||||
// cleanupOldExecutions 清理旧的执行记录(保持内存中的记录数量在限制内)
|
||||
func (m *ExternalMCPManager) cleanupOldExecutions() {
|
||||
const maxExecutionsInMemory = 1000
|
||||
if len(m.executions) <= maxExecutionsInMemory {
|
||||
return
|
||||
}
|
||||
|
||||
// 按开始时间排序,删除最旧的记录
|
||||
type execTime struct {
|
||||
id string
|
||||
startTime time.Time
|
||||
}
|
||||
var execs []execTime
|
||||
for id, exec := range m.executions {
|
||||
execs = append(execs, execTime{id: id, startTime: exec.StartTime})
|
||||
}
|
||||
|
||||
// 按时间排序
|
||||
for i := 0; i < len(execs)-1; i++ {
|
||||
for j := i + 1; j < len(execs); j++ {
|
||||
if execs[i].startTime.After(execs[j].startTime) {
|
||||
execs[i], execs[j] = execs[j], execs[i]
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 删除最旧的记录
|
||||
toDelete := len(m.executions) - maxExecutionsInMemory
|
||||
for i := 0; i < toDelete && i < len(execs); i++ {
|
||||
delete(m.executions, execs[i].id)
|
||||
}
|
||||
}
|
||||
|
||||
// GetExecution 获取执行记录(先从内存查找,再从数据库查找)
|
||||
func (m *ExternalMCPManager) GetExecution(id string) (*ToolExecution, bool) {
|
||||
m.mu.RLock()
|
||||
exec, exists := m.executions[id]
|
||||
m.mu.RUnlock()
|
||||
|
||||
if exists {
|
||||
return exec, true
|
||||
}
|
||||
|
||||
if m.storage != nil {
|
||||
exec, err := m.storage.GetToolExecution(id)
|
||||
if err == nil {
|
||||
return exec, true
|
||||
}
|
||||
}
|
||||
|
||||
return nil, false
|
||||
}
|
||||
|
||||
// updateStats 更新统计信息
|
||||
func (m *ExternalMCPManager) updateStats(toolName string, failed bool) {
|
||||
now := time.Now()
|
||||
if m.storage != nil {
|
||||
totalCalls := 1
|
||||
successCalls := 0
|
||||
failedCalls := 0
|
||||
if failed {
|
||||
failedCalls = 1
|
||||
} else {
|
||||
successCalls = 1
|
||||
}
|
||||
if err := m.storage.UpdateToolStats(toolName, totalCalls, successCalls, failedCalls, &now); err != nil {
|
||||
m.logger.Warn("保存统计信息到数据库失败", zap.Error(err))
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
if m.stats[toolName] == nil {
|
||||
m.stats[toolName] = &ToolStats{
|
||||
ToolName: toolName,
|
||||
}
|
||||
}
|
||||
|
||||
stats := m.stats[toolName]
|
||||
stats.TotalCalls++
|
||||
stats.LastCallTime = &now
|
||||
|
||||
if failed {
|
||||
stats.FailedCalls++
|
||||
} else {
|
||||
stats.SuccessCalls++
|
||||
}
|
||||
}
|
||||
|
||||
// GetStats 获取MCP服务器统计信息
|
||||
func (m *ExternalMCPManager) GetStats() map[string]interface{} {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
|
||||
total := len(m.configs)
|
||||
enabled := 0
|
||||
disabled := 0
|
||||
connected := 0
|
||||
|
||||
for name, cfg := range m.configs {
|
||||
if m.isEnabled(cfg) {
|
||||
enabled++
|
||||
if client, exists := m.clients[name]; exists && client.IsConnected() {
|
||||
connected++
|
||||
}
|
||||
} else {
|
||||
disabled++
|
||||
}
|
||||
}
|
||||
|
||||
return map[string]interface{}{
|
||||
"total": total,
|
||||
"enabled": enabled,
|
||||
"disabled": disabled,
|
||||
"connected": connected,
|
||||
}
|
||||
}
|
||||
|
||||
// GetToolStats 获取工具统计信息(合并内存和数据库)
|
||||
// 只返回外部MCP工具的统计信息(工具名称包含 "::")
|
||||
func (m *ExternalMCPManager) GetToolStats() map[string]*ToolStats {
|
||||
result := make(map[string]*ToolStats)
|
||||
|
||||
// 从数据库加载统计信息(如果使用数据库存储)
|
||||
if m.storage != nil {
|
||||
dbStats, err := m.storage.LoadToolStats()
|
||||
if err == nil {
|
||||
// 只保留外部MCP工具的统计信息(工具名称包含 "::")
|
||||
for k, v := range dbStats {
|
||||
if findSubstring(k, "::") > 0 {
|
||||
result[k] = v
|
||||
}
|
||||
}
|
||||
} else {
|
||||
m.logger.Warn("从数据库加载统计信息失败", zap.Error(err))
|
||||
}
|
||||
}
|
||||
|
||||
// 合并内存中的统计信息
|
||||
m.mu.RLock()
|
||||
for k, v := range m.stats {
|
||||
// 如果数据库中已有该工具的统计信息,合并它们
|
||||
if existing, exists := result[k]; exists {
|
||||
// 创建新的统计信息对象,避免修改共享对象
|
||||
merged := &ToolStats{
|
||||
ToolName: k,
|
||||
TotalCalls: existing.TotalCalls + v.TotalCalls,
|
||||
SuccessCalls: existing.SuccessCalls + v.SuccessCalls,
|
||||
FailedCalls: existing.FailedCalls + v.FailedCalls,
|
||||
}
|
||||
// 使用最新的调用时间
|
||||
if v.LastCallTime != nil && (existing.LastCallTime == nil || v.LastCallTime.After(*existing.LastCallTime)) {
|
||||
merged.LastCallTime = v.LastCallTime
|
||||
} else if existing.LastCallTime != nil {
|
||||
timeCopy := *existing.LastCallTime
|
||||
merged.LastCallTime = &timeCopy
|
||||
}
|
||||
result[k] = merged
|
||||
} else {
|
||||
// 如果数据库中没有,直接使用内存中的统计信息
|
||||
statCopy := *v
|
||||
result[k] = &statCopy
|
||||
}
|
||||
}
|
||||
m.mu.RUnlock()
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// GetToolCount 获取指定外部MCP的工具数量
|
||||
func (m *ExternalMCPManager) GetToolCount(name string) (int, error) {
|
||||
client, exists := m.GetClient(name)
|
||||
if !exists {
|
||||
return 0, fmt.Errorf("客户端不存在: %s", name)
|
||||
}
|
||||
|
||||
if !client.IsConnected() {
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
|
||||
tools, err := client.ListTools(ctx)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("获取工具列表失败: %w", err)
|
||||
}
|
||||
|
||||
return len(tools), nil
|
||||
}
|
||||
|
||||
// GetToolCounts 获取所有外部MCP的工具数量
|
||||
func (m *ExternalMCPManager) GetToolCounts() map[string]int {
|
||||
m.mu.RLock()
|
||||
clients := make(map[string]ExternalMCPClient)
|
||||
for k, v := range m.clients {
|
||||
clients[k] = v
|
||||
}
|
||||
m.mu.RUnlock()
|
||||
|
||||
result := make(map[string]int)
|
||||
for name, client := range clients {
|
||||
if !client.IsConnected() {
|
||||
result[name] = 0
|
||||
continue
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
tools, err := client.ListTools(ctx)
|
||||
cancel()
|
||||
|
||||
if err != nil {
|
||||
m.logger.Warn("获取外部MCP工具数量失败",
|
||||
zap.String("name", name),
|
||||
zap.Error(err),
|
||||
)
|
||||
result[name] = 0
|
||||
continue
|
||||
}
|
||||
|
||||
result[name] = len(tools)
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// connectClient 连接客户端(异步)
|
||||
func (m *ExternalMCPManager) connectClient(name string, serverCfg config.ExternalMCPServerConfig) error {
|
||||
var client ExternalMCPClient
|
||||
|
||||
timeout := time.Duration(serverCfg.Timeout) * time.Second
|
||||
if timeout <= 0 {
|
||||
timeout = 30 * time.Second
|
||||
}
|
||||
|
||||
// 根据传输模式创建客户端
|
||||
transport := serverCfg.Transport
|
||||
if transport == "" {
|
||||
// 如果没有指定transport,根据是否有command或url判断
|
||||
if serverCfg.Command != "" {
|
||||
transport = "stdio"
|
||||
} else if serverCfg.URL != "" {
|
||||
transport = "http"
|
||||
} else {
|
||||
return fmt.Errorf("无法确定传输模式: 需要指定command或url")
|
||||
}
|
||||
}
|
||||
|
||||
switch transport {
|
||||
case "http":
|
||||
if serverCfg.URL == "" {
|
||||
return fmt.Errorf("HTTP模式需要URL")
|
||||
}
|
||||
client = NewHTTPMCPClient(serverCfg.URL, timeout, m.logger)
|
||||
case "stdio":
|
||||
if serverCfg.Command == "" {
|
||||
return fmt.Errorf("stdio模式需要command")
|
||||
}
|
||||
client = NewStdioMCPClient(serverCfg.Command, serverCfg.Args, timeout, m.logger)
|
||||
default:
|
||||
return fmt.Errorf("不支持的传输模式: %s", transport)
|
||||
}
|
||||
|
||||
// 初始化连接
|
||||
ctx, cancel := context.WithTimeout(context.Background(), timeout)
|
||||
defer cancel()
|
||||
|
||||
if err := client.Initialize(ctx); err != nil {
|
||||
m.logger.Error("初始化外部MCP客户端失败",
|
||||
zap.String("name", name),
|
||||
zap.Error(err),
|
||||
)
|
||||
return err
|
||||
}
|
||||
|
||||
// 保存客户端
|
||||
m.mu.Lock()
|
||||
m.clients[name] = client
|
||||
m.mu.Unlock()
|
||||
|
||||
m.logger.Info("外部MCP客户端已连接",
|
||||
zap.String("name", name),
|
||||
zap.String("transport", transport),
|
||||
)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// isEnabled 检查是否启用
|
||||
func (m *ExternalMCPManager) isEnabled(cfg config.ExternalMCPServerConfig) bool {
|
||||
// 优先使用 ExternalMCPEnable 字段
|
||||
// 如果没有设置,检查旧的 enabled/disabled 字段(向后兼容)
|
||||
if cfg.ExternalMCPEnable {
|
||||
return true
|
||||
}
|
||||
// 向后兼容:检查旧字段
|
||||
if cfg.Disabled {
|
||||
return false
|
||||
}
|
||||
if cfg.Enabled {
|
||||
return true
|
||||
}
|
||||
// 都没有设置,默认为启用
|
||||
return true
|
||||
}
|
||||
|
||||
// findSubstring 查找子字符串(简单实现)
|
||||
func findSubstring(s, substr string) int {
|
||||
for i := 0; i <= len(s)-len(substr); i++ {
|
||||
if s[i:i+len(substr)] == substr {
|
||||
return i
|
||||
}
|
||||
}
|
||||
return -1
|
||||
}
|
||||
|
||||
// StartAllEnabled 启动所有启用的客户端
|
||||
func (m *ExternalMCPManager) StartAllEnabled() {
|
||||
m.mu.RLock()
|
||||
configs := make(map[string]config.ExternalMCPServerConfig)
|
||||
for k, v := range m.configs {
|
||||
configs[k] = v
|
||||
}
|
||||
m.mu.RUnlock()
|
||||
|
||||
for name, cfg := range configs {
|
||||
if m.isEnabled(cfg) {
|
||||
go func(n string, c config.ExternalMCPServerConfig) {
|
||||
if err := m.connectClient(n, c); err != nil {
|
||||
m.logger.Error("启动外部MCP客户端失败",
|
||||
zap.String("name", n),
|
||||
zap.Error(err),
|
||||
)
|
||||
}
|
||||
}(name, cfg)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// StopAll 停止所有客户端
|
||||
func (m *ExternalMCPManager) StopAll() {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
for name, client := range m.clients {
|
||||
client.Close()
|
||||
delete(m.clients, name)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,261 @@
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"cyberstrike-ai/internal/config"
|
||||
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
func TestExternalMCPManager_AddOrUpdateConfig(t *testing.T) {
|
||||
logger := zap.NewNop()
|
||||
manager := NewExternalMCPManager(logger)
|
||||
|
||||
// 测试添加stdio配置
|
||||
stdioCfg := config.ExternalMCPServerConfig{
|
||||
Command: "python3",
|
||||
Args: []string{"/path/to/script.py"},
|
||||
Transport: "stdio",
|
||||
Description: "Test stdio MCP",
|
||||
Timeout: 30,
|
||||
Enabled: true,
|
||||
}
|
||||
|
||||
err := manager.AddOrUpdateConfig("test-stdio", stdioCfg)
|
||||
if err != nil {
|
||||
t.Fatalf("添加stdio配置失败: %v", err)
|
||||
}
|
||||
|
||||
// 测试添加HTTP配置
|
||||
httpCfg := config.ExternalMCPServerConfig{
|
||||
Transport: "http",
|
||||
URL: "http://127.0.0.1:8081/mcp",
|
||||
Description: "Test HTTP MCP",
|
||||
Timeout: 30,
|
||||
Enabled: false,
|
||||
}
|
||||
|
||||
err = manager.AddOrUpdateConfig("test-http", httpCfg)
|
||||
if err != nil {
|
||||
t.Fatalf("添加HTTP配置失败: %v", err)
|
||||
}
|
||||
|
||||
// 验证配置已保存
|
||||
configs := manager.GetConfigs()
|
||||
if len(configs) != 2 {
|
||||
t.Fatalf("期望2个配置,实际%d个", len(configs))
|
||||
}
|
||||
|
||||
if configs["test-stdio"].Command != stdioCfg.Command {
|
||||
t.Errorf("stdio配置命令不匹配")
|
||||
}
|
||||
|
||||
if configs["test-http"].URL != httpCfg.URL {
|
||||
t.Errorf("HTTP配置URL不匹配")
|
||||
}
|
||||
}
|
||||
|
||||
func TestExternalMCPManager_RemoveConfig(t *testing.T) {
|
||||
logger := zap.NewNop()
|
||||
manager := NewExternalMCPManager(logger)
|
||||
|
||||
cfg := config.ExternalMCPServerConfig{
|
||||
Command: "python3",
|
||||
Transport: "stdio",
|
||||
Enabled: false,
|
||||
}
|
||||
|
||||
manager.AddOrUpdateConfig("test-remove", cfg)
|
||||
|
||||
// 移除配置
|
||||
err := manager.RemoveConfig("test-remove")
|
||||
if err != nil {
|
||||
t.Fatalf("移除配置失败: %v", err)
|
||||
}
|
||||
|
||||
configs := manager.GetConfigs()
|
||||
if _, exists := configs["test-remove"]; exists {
|
||||
t.Error("配置应该已被移除")
|
||||
}
|
||||
}
|
||||
|
||||
func TestExternalMCPManager_GetStats(t *testing.T) {
|
||||
logger := zap.NewNop()
|
||||
manager := NewExternalMCPManager(logger)
|
||||
|
||||
// 添加多个配置
|
||||
manager.AddOrUpdateConfig("enabled1", config.ExternalMCPServerConfig{
|
||||
Command: "python3",
|
||||
Enabled: true,
|
||||
})
|
||||
|
||||
manager.AddOrUpdateConfig("enabled2", config.ExternalMCPServerConfig{
|
||||
URL: "http://127.0.0.1:8081/mcp",
|
||||
Enabled: true,
|
||||
})
|
||||
|
||||
manager.AddOrUpdateConfig("disabled1", config.ExternalMCPServerConfig{
|
||||
Command: "python3",
|
||||
Enabled: false,
|
||||
Disabled: true, // 明确设置为禁用
|
||||
})
|
||||
|
||||
stats := manager.GetStats()
|
||||
|
||||
if stats["total"].(int) != 3 {
|
||||
t.Errorf("期望总数3,实际%d", stats["total"])
|
||||
}
|
||||
|
||||
if stats["enabled"].(int) != 2 {
|
||||
t.Errorf("期望启用数2,实际%d", stats["enabled"])
|
||||
}
|
||||
|
||||
if stats["disabled"].(int) != 1 {
|
||||
t.Errorf("期望停用数1,实际%d", stats["disabled"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestExternalMCPManager_LoadConfigs(t *testing.T) {
|
||||
logger := zap.NewNop()
|
||||
manager := NewExternalMCPManager(logger)
|
||||
|
||||
externalMCPConfig := config.ExternalMCPConfig{
|
||||
Servers: map[string]config.ExternalMCPServerConfig{
|
||||
"loaded1": {
|
||||
Command: "python3",
|
||||
Enabled: true,
|
||||
},
|
||||
"loaded2": {
|
||||
URL: "http://127.0.0.1:8081/mcp",
|
||||
Enabled: false,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
manager.LoadConfigs(&externalMCPConfig)
|
||||
|
||||
configs := manager.GetConfigs()
|
||||
if len(configs) != 2 {
|
||||
t.Fatalf("期望2个配置,实际%d个", len(configs))
|
||||
}
|
||||
|
||||
if configs["loaded1"].Command != "python3" {
|
||||
t.Error("配置1加载失败")
|
||||
}
|
||||
|
||||
if configs["loaded2"].URL != "http://127.0.0.1:8081/mcp" {
|
||||
t.Error("配置2加载失败")
|
||||
}
|
||||
}
|
||||
|
||||
func TestHTTPMCPClient_Initialize(t *testing.T) {
|
||||
// 注意:这个测试需要一个真实的HTTP MCP服务器
|
||||
// 如果没有服务器,这个测试会失败
|
||||
// 在实际测试中,可以使用mock服务器
|
||||
logger := zap.NewNop()
|
||||
client := NewHTTPMCPClient("http://127.0.0.1:8081/mcp", 5*time.Second, logger)
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
|
||||
// 这个测试可能会失败,如果没有真实的服务器
|
||||
// 在实际环境中,应该使用mock服务器
|
||||
err := client.Initialize(ctx)
|
||||
if err != nil {
|
||||
t.Logf("初始化失败(可能是没有服务器): %v", err)
|
||||
}
|
||||
|
||||
status := client.GetStatus()
|
||||
if status == "" {
|
||||
t.Error("状态不应该为空")
|
||||
}
|
||||
|
||||
client.Close()
|
||||
}
|
||||
|
||||
func TestStdioMCPClient_Initialize(t *testing.T) {
|
||||
// 注意:这个测试需要一个真实的stdio MCP服务器
|
||||
// 如果没有服务器,这个测试会失败
|
||||
logger := zap.NewNop()
|
||||
client := NewStdioMCPClient("echo", []string{"test"}, 5*time.Second, logger)
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
|
||||
// 这个测试可能会失败,因为echo不是MCP服务器
|
||||
// 在实际环境中,应该使用真实的MCP服务器或mock
|
||||
err := client.Initialize(ctx)
|
||||
if err != nil {
|
||||
t.Logf("初始化失败(echo不是MCP服务器): %v", err)
|
||||
}
|
||||
|
||||
client.Close()
|
||||
}
|
||||
|
||||
func TestExternalMCPManager_StartStopClient(t *testing.T) {
|
||||
logger := zap.NewNop()
|
||||
manager := NewExternalMCPManager(logger)
|
||||
|
||||
// 添加一个禁用的配置
|
||||
cfg := config.ExternalMCPServerConfig{
|
||||
Command: "python3",
|
||||
Transport: "stdio",
|
||||
Enabled: false,
|
||||
}
|
||||
|
||||
manager.AddOrUpdateConfig("test-start-stop", cfg)
|
||||
|
||||
// 尝试启动(可能会失败,因为没有真实的服务器)
|
||||
err := manager.StartClient("test-start-stop")
|
||||
if err != nil {
|
||||
t.Logf("启动失败(可能是没有服务器): %v", err)
|
||||
}
|
||||
|
||||
// 停止
|
||||
err = manager.StopClient("test-start-stop")
|
||||
if err != nil {
|
||||
t.Fatalf("停止失败: %v", err)
|
||||
}
|
||||
|
||||
// 验证配置已更新为禁用
|
||||
configs := manager.GetConfigs()
|
||||
if configs["test-start-stop"].Enabled {
|
||||
t.Error("配置应该已被禁用")
|
||||
}
|
||||
}
|
||||
|
||||
func TestExternalMCPManager_CallTool(t *testing.T) {
|
||||
logger := zap.NewNop()
|
||||
manager := NewExternalMCPManager(logger)
|
||||
|
||||
// 测试调用不存在的工具
|
||||
_, err := manager.CallTool(context.Background(), "nonexistent::tool", map[string]interface{}{})
|
||||
if err == nil {
|
||||
t.Error("应该返回错误")
|
||||
}
|
||||
|
||||
// 测试无效的工具名称格式
|
||||
_, err = manager.CallTool(context.Background(), "invalid-tool-name", map[string]interface{}{})
|
||||
if err == nil {
|
||||
t.Error("应该返回错误(无效格式)")
|
||||
}
|
||||
}
|
||||
|
||||
func TestExternalMCPManager_GetAllTools(t *testing.T) {
|
||||
logger := zap.NewNop()
|
||||
manager := NewExternalMCPManager(logger)
|
||||
|
||||
ctx := context.Background()
|
||||
tools, err := manager.GetAllTools(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("获取工具列表失败: %v", err)
|
||||
}
|
||||
|
||||
// 如果没有连接的客户端,应该返回空列表
|
||||
if len(tools) != 0 {
|
||||
t.Logf("获取到%d个工具", len(tools))
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user