refactor: split agent.go into 12 files + add mcp_restart_server tool

This commit is contained in:
root
2026-07-22 15:40:27 +08:00
parent 76a4619ed5
commit 85992902d9
20 changed files with 3150 additions and 3080 deletions

View File

@ -134,3 +134,11 @@ func (s *Server) CallTool(name string, args map[string]interface{}) (string, err
func (s *Server) Close() error {
return s.transport.Close()
}
func (s *Server) SetTransport(t Transport) {
s.mu.Lock()
defer s.mu.Unlock()
s.transport.Close()
s.transport = t
s.nextID = 0
}

View File

@ -30,6 +30,8 @@ func init() {
type Plugin struct {
name string
servers []*Server
configs []serverConfig
sdk *sdk.PluginSDK
mu sync.Mutex
wg sync.WaitGroup
}
@ -42,41 +44,39 @@ func (p *Plugin) Name() string { return p.name }
func (p *Plugin) Start(s *sdk.PluginSDK) error {
s.SetAutoRestart(true)
p.sdk = s
cfgs, err := p.loadConfig(s)
if err != nil {
return fmt.Errorf("load mcp config: %w", err)
}
p.configs = cfgs
if len(cfgs) == 0 {
log.Printf("[mcp] no servers configured, idle")
return nil
}
for _, cfg := range cfgs {
server, tools, err := p.connectServer(cfg)
if err != nil {
if err := p.connectAndRegister(cfg); err != nil {
log.Printf("[mcp] connect %s: %v", cfg.Name, err)
continue
}
}
for _, tool := range tools {
toolName := fmt.Sprintf("%s_%s", cfg.Name, tool.Name)
tDef := sdk.ToolDef{
Name: toolName,
Description: fmt.Sprintf("[MCP/%s] %s", cfg.Name, tool.Description),
Parameters: tool.InputSchema,
}
tHandler := p.makeHandler(server, tool.Name)
if err := s.RegisterTool(toolName, tDef, tHandler); err != nil {
log.Printf("[mcp] register tool %s: %v", toolName, err)
continue
}
log.Printf("[mcp] registered tool: %s (%s)", toolName, cfg.Name)
}
p.mu.Lock()
p.servers = append(p.servers, server)
p.mu.Unlock()
log.Printf("[mcp] connected server: %s (%d tools)", cfg.Name, len(tools))
tDef := sdk.ToolDef{
Name: "mcp_restart_server",
Description: "重启 MCP 服务器连接。当 MCP 工具返回 pipe/transport closed 错误时,用此工具重启指定的 MCP 服务器。",
Parameters: map[string]interface{}{
"type": "object",
"properties": map[string]interface{}{
"name": map[string]interface{}{
"type": "string",
"description": "MCP 服务器名称(如 email",
},
},
"required": []string{"name"},
},
}
if err := s.RegisterTool("mcp_restart_server", tDef, p.restartServerHandler); err != nil {
log.Printf("[mcp] register restart tool: %v", err)
}
return nil
@ -178,6 +178,91 @@ func (p *Plugin) makeHandler(server *Server, toolName string) sdk.ToolHandler {
}
}
func (p *Plugin) connectAndRegister(cfg serverConfig) error {
server, tools, err := p.connectServer(cfg)
if err != nil {
return err
}
for _, tool := range tools {
toolName := fmt.Sprintf("%s_%s", cfg.Name, tool.Name)
tDef := sdk.ToolDef{
Name: toolName,
Description: fmt.Sprintf("[MCP/%s] %s", cfg.Name, tool.Description),
Parameters: tool.InputSchema,
}
tHandler := p.makeHandler(server, tool.Name)
if err := p.sdk.RegisterTool(toolName, tDef, tHandler); err != nil {
log.Printf("[mcp] register tool %s: %v", toolName, err)
continue
}
log.Printf("[mcp] registered tool: %s (%s)", toolName, cfg.Name)
}
p.mu.Lock()
p.servers = append(p.servers, server)
p.mu.Unlock()
log.Printf("[mcp] connected server: %s (%d tools)", cfg.Name, len(tools))
return nil
}
func (p *Plugin) restartServerHandler(args map[string]interface{}) (interface{}, error) {
name, _ := args["name"].(string)
if name == "" {
return "参数 name 不能为空", nil
}
p.mu.Lock()
var idx int = -1
for i, s := range p.servers {
if s.Name() == name {
idx = i
break
}
}
if idx == -1 {
p.mu.Unlock()
return fmt.Sprintf("MCP 服务器 [%s] 不存在", name), nil
}
oldServer := p.servers[idx]
p.mu.Unlock()
var cfg *serverConfig
for i := range p.configs {
if p.configs[i].Name == name {
cfg = &p.configs[i]
break
}
}
if cfg == nil {
return fmt.Sprintf("MCP 服务器 [%s] 的配置未找到", name), nil
}
transport, err := newTransport(*cfg)
if err != nil {
return fmt.Sprintf("创建 MCP 服务器 [%s] 传输层失败: %v", name, err), nil
}
oldServer.SetTransport(transport)
_, err = oldServer.ListTools()
if err != nil {
return fmt.Sprintf("MCP 服务器 [%s] 重启后通信仍异常: %v", name, err), nil
}
return fmt.Sprintf("MCP 服务器 [%s] 已成功重启", name), nil
}
func newTransport(cfg serverConfig) (Transport, error) {
if cfg.URL != "" {
return NewSSETransport(cfg.URL), nil
}
if cfg.Command != "" {
return NewStdioTransport(cfg.Command, cfg.Args, cfg.Env)
}
return nil, fmt.Errorf("neither command nor url specified")
}
func (p *Plugin) Stop() error {
p.mu.Lock()
defer p.mu.Unlock()