mirror of
https://gitcode.com/JianFeeeee/HomeAgent.git
synced 2026-09-22 18:08:04 +00:00
refactor: split agent.go into 12 files + add mcp_restart_server tool
This commit is contained in:
@ -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
|
||||
}
|
||||
|
||||
@ -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()
|
||||
|
||||
Reference in New Issue
Block a user