Files
HomeAgent/internal/plugins/mcp/plugin.go

189 lines
4.6 KiB
Go
Raw 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 mcp
import (
"encoding/json"
"fmt"
"log"
"strings"
"sync"
"gitcode.com/JianFeeeee/HomeAgent/internal/plugin"
sdk "gitcode.com/JianFeeeee/HomeAgent/internal/sdk"
)
// MCP 服务器配置(来自 config_mcp 表或 skill.json
type serverConfig struct {
Name string `json:"name"`
Command string `json:"command,omitempty"` // stdio 模式
Args []string `json:"args,omitempty"`
Env []string `json:"env,omitempty"`
URL string `json:"url,omitempty"` // SSE 模式
}
func init() {
plugin.RegisterPluginMeta("mcp", "MCP 服务器", "MCP")
plugin.RegisterFactory("mcp", func(name string, config map[string]interface{}) (sdk.Plugin, error) {
return New(name), nil
})
}
type Plugin struct {
name string
servers []*Server
mu sync.Mutex
wg sync.WaitGroup
}
func New(name string) *Plugin {
return &Plugin{name: name}
}
func (p *Plugin) Name() string { return p.name }
func (p *Plugin) Start(s *sdk.PluginSDK) error {
cfgs, err := p.loadConfig(s)
if err != nil {
return fmt.Errorf("load mcp config: %w", err)
}
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 {
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))
}
return nil
}
func (p *Plugin) loadConfig(s *sdk.PluginSDK) ([]serverConfig, error) {
// 优先从独立服务器配置键读取servers.<name>.<field>
keys, _ := s.Settings().List("servers.")
if len(keys) > 0 {
serverNames := make(map[string]bool)
for _, k := range keys {
parts := strings.SplitN(k, ".", 3)
if len(parts) >= 2 {
serverNames[parts[1]] = true
}
}
var cfgs []serverConfig
for name := range serverNames {
cfg := serverConfig{Name: name}
if v, _ := s.Settings().Get("servers." + name + ".command"); v != nil {
if s, ok := v.(string); ok {
cfg.Command = s
}
}
if v, _ := s.Settings().Get("servers." + name + ".url"); v != nil {
if s, ok := v.(string); ok {
cfg.URL = s
}
}
if v, _ := s.Settings().Get("servers." + name + ".args"); v != nil {
if s, ok := v.(string); ok && s != "" {
json.Unmarshal([]byte(s), &cfg.Args)
}
}
if v, _ := s.Settings().Get("servers." + name + ".env"); v != nil {
if s, ok := v.(string); ok && s != "" {
json.Unmarshal([]byte(s), &cfg.Env)
}
}
if cfg.Command != "" || cfg.URL != "" {
cfgs = append(cfgs, cfg)
}
}
if len(cfgs) > 0 {
return cfgs, nil
}
}
// 回退:从旧版 JSON blob 读取
raw, err := s.Settings().Get("servers")
if err == nil {
switch v := raw.(type) {
case string:
var cfgs []serverConfig
if err := json.Unmarshal([]byte(v), &cfgs); err == nil && len(cfgs) > 0 {
return cfgs, nil
}
case []interface{}:
data, _ := json.Marshal(v)
var cfgs []serverConfig
if json.Unmarshal(data, &cfgs) == nil && len(cfgs) > 0 {
return cfgs, nil
}
}
}
return nil, nil
}
func (p *Plugin) connectServer(cfg serverConfig) (*Server, []MCPTool, error) {
var transport Transport
if cfg.URL != "" {
transport = NewSSETransport(cfg.URL)
} else if cfg.Command != "" {
var err error
transport, err = NewStdioTransport(cfg.Command, cfg.Args, cfg.Env)
if err != nil {
return nil, nil, fmt.Errorf("stdio transport: %w", err)
}
} else {
return nil, nil, fmt.Errorf("neither command nor url specified")
}
server := NewServer(cfg.Name, transport)
tools, err := server.ListTools()
if err != nil {
transport.Close()
return nil, nil, fmt.Errorf("list tools: %w", err)
}
return server, tools, nil
}
func (p *Plugin) makeHandler(server *Server, toolName string) sdk.ToolHandler {
return func(args map[string]interface{}) (interface{}, error) {
return server.CallTool(toolName, args)
}
}
func (p *Plugin) Stop() error {
p.mu.Lock()
defer p.mu.Unlock()
for _, s := range p.servers {
s.Close()
}
p.servers = nil
return nil
}