feat(mcp): implement server connection management and transport handling

This commit is contained in:
tong3jie 2026-03-17 12:48:35 +08:00
parent cff85cfe5c
commit 651fe4f04d
5 changed files with 461 additions and 334 deletions

1
go.mod
View file

@ -80,6 +80,7 @@ require (
github.com/gogo/protobuf v1.3.2 // indirect github.com/gogo/protobuf v1.3.2 // indirect
github.com/google/jsonschema-go v0.4.2 // indirect github.com/google/jsonschema-go v0.4.2 // indirect
github.com/grbit/go-json v0.11.0 // indirect github.com/grbit/go-json v0.11.0 // indirect
github.com/joho/godotenv v1.5.1
github.com/klauspost/compress v1.18.4 // indirect github.com/klauspost/compress v1.18.4 // indirect
github.com/klauspost/cpuid/v2 v2.3.0 // indirect github.com/klauspost/cpuid/v2 v2.3.0 // indirect
github.com/tidwall/gjson v1.18.0 // indirect github.com/tidwall/gjson v1.18.0 // indirect

2
go.sum
View file

@ -111,6 +111,8 @@ github.com/hashicorp/golang-lru/v2 v2.0.7/go.mod h1:QeFd9opnmA6QUJc5vARoKUSoFhyf
github.com/hpcloud/tail v1.0.0/go.mod h1:ab1qPbhIpdTxEkNHXyeSf5vhxWSCs/tWer42PpOxQnU= github.com/hpcloud/tail v1.0.0/go.mod h1:ab1qPbhIpdTxEkNHXyeSf5vhxWSCs/tWer42PpOxQnU=
github.com/inconshreveable/mousetrap v1.1.0 h1:wN+x4NVGpMsO7ErUn/mUI3vEoE6Jt13X2s0bqwp9tc8= github.com/inconshreveable/mousetrap v1.1.0 h1:wN+x4NVGpMsO7ErUn/mUI3vEoE6Jt13X2s0bqwp9tc8=
github.com/inconshreveable/mousetrap v1.1.0/go.mod h1:vpF70FUmC8bwa3OWnCshd2FqLfsEA9PFc4w1p2J65bw= github.com/inconshreveable/mousetrap v1.1.0/go.mod h1:vpF70FUmC8bwa3OWnCshd2FqLfsEA9PFc4w1p2J65bw=
github.com/joho/godotenv v1.5.1 h1:7eLL/+HRGLY0ldzfGMeQkb7vMd0as4CfYvUVzLqw0N0=
github.com/joho/godotenv v1.5.1/go.mod h1:f4LDr5Voq0i2e/R5DDNOoa2zzDfwtkZa6DnEwAbqwq4=
github.com/kisielk/errcheck v1.5.0/go.mod h1:pFxgyoBC7bSaBwPgfKdkLd5X25qrDl4LWUI2bnpBCr8= github.com/kisielk/errcheck v1.5.0/go.mod h1:pFxgyoBC7bSaBwPgfKdkLd5X25qrDl4LWUI2bnpBCr8=
github.com/kisielk/gotool v1.0.0/go.mod h1:XhKaO+MFFWcvkIS/tQcRk01m1F5IRFswLeQ+oQHNcck= github.com/kisielk/gotool v1.0.0/go.mod h1:XhKaO+MFFWcvkIS/tQcRk01m1F5IRFswLeQ+oQHNcck=
github.com/klauspost/compress v1.18.4 h1:RPhnKRAQ4Fh8zU2FY/6ZFDwTVTxgJ/EMydqSTzE9a2c= github.com/klauspost/compress v1.18.4 h1:RPhnKRAQ4Fh8zU2FY/6ZFDwTVTxgJ/EMydqSTzE9a2c=

View file

@ -1,17 +1,16 @@
package mcp package mcp
import ( import (
"bufio" "cmp"
"context" "context"
"errors" "errors"
"fmt" "fmt"
"maps"
"net/http" "net/http"
"os"
"os/exec"
"path/filepath" "path/filepath"
"strings"
"sync" "sync"
"sync/atomic" "sync/atomic"
"time"
"github.com/modelcontextprotocol/go-sdk/mcp" "github.com/modelcontextprotocol/go-sdk/mcp"
@ -19,6 +18,22 @@ import (
"github.com/sipeed/picoclaw/pkg/logger" "github.com/sipeed/picoclaw/pkg/logger"
) )
const (
clientName = "picoclaw"
clientVersion = "1.0.0"
logModule = "mcp"
maxFails = 3
)
var (
ErrManagerClosed = errors.New("MCP manager is closed")
ErrInvalidServerConfig = errors.New("either URL or command must be provided")
ErrStdioCommandRequired = errors.New("command is required for stdio transport")
)
// headerTransport is an http.RoundTripper that adds custom headers to requests // headerTransport is an http.RoundTripper that adds custom headers to requests
type headerTransport struct { type headerTransport struct {
base http.RoundTripper base http.RoundTripper
@ -26,91 +41,21 @@ type headerTransport struct {
} }
func (t *headerTransport) RoundTrip(req *http.Request) (*http.Response, error) { func (t *headerTransport) RoundTrip(req *http.Request) (*http.Response, error) {
// Clone the request to avoid modifying the original
req = req.Clone(req.Context()) req = req.Clone(req.Context())
// Add custom headers
for key, value := range t.headers { for key, value := range t.headers {
req.Header.Set(key, value) req.Header.Add(key, value)
} }
base := cmp.Or(t.base, http.DefaultTransport)
// Use the base transport
base := t.base
if base == nil {
base = http.DefaultTransport
}
return base.RoundTrip(req) return base.RoundTrip(req)
} }
// loadEnvFile loads environment variables from a file in .env format
// Each line should be in the format: KEY=value
// Lines starting with # are comments
// Empty lines are ignored
func loadEnvFile(path string) (map[string]string, error) {
file, err := os.Open(path)
if err != nil {
return nil, fmt.Errorf("failed to open env file: %w", err)
}
defer file.Close()
envVars := make(map[string]string)
scanner := bufio.NewScanner(file)
lineNum := 0
for scanner.Scan() {
lineNum++
line := strings.TrimSpace(scanner.Text())
// Skip empty lines and comments
if line == "" || strings.HasPrefix(line, "#") {
continue
}
// Parse KEY=value
parts := strings.SplitN(line, "=", 2)
if len(parts) != 2 {
return nil, fmt.Errorf("invalid format at line %d: %s", lineNum, line)
}
key := strings.TrimSpace(parts[0])
value := strings.TrimSpace(parts[1])
if key == "" {
return nil, fmt.Errorf("invalid format at line %d: empty key", lineNum)
}
// Remove surrounding quotes if present
if len(value) >= 2 {
if (value[0] == '"' && value[len(value)-1] == '"') ||
(value[0] == '\'' && value[len(value)-1] == '\'') {
value = value[1 : len(value)-1]
}
}
envVars[key] = value
}
if err := scanner.Err(); err != nil {
return nil, fmt.Errorf("error reading env file: %w", err)
}
return envVars, nil
}
// ServerConnection represents a connection to an MCP server
type ServerConnection struct {
Name string
Client *mcp.Client
Session *mcp.ClientSession
Tools []*mcp.Tool
}
// Manager manages multiple MCP server connections // Manager manages multiple MCP server connections
type Manager struct { type Manager struct {
servers map[string]*ServerConnection servers map[string]*ServerConnection
mu sync.RWMutex mu sync.RWMutex
closed atomic.Bool // changed from bool to atomic.Bool to avoid TOCTOU race closed atomic.Bool
wg sync.WaitGroup // tracks in-flight CallTool calls wg sync.WaitGroup
} }
// NewManager creates a new MCP manager // NewManager creates a new MCP manager
@ -133,19 +78,18 @@ func (m *Manager) LoadFromMCPConfig(
workspacePath string, workspacePath string,
) error { ) error {
if !mcpCfg.Enabled { if !mcpCfg.Enabled {
logger.InfoCF("mcp", "MCP integration is disabled", nil) logger.InfoCF(logModule, "MCP integration is disabled", nil)
return nil return nil
} }
if len(mcpCfg.Servers) == 0 { if len(mcpCfg.Servers) == 0 {
logger.InfoCF("mcp", "No MCP servers configured", nil) logger.InfoCF(logModule, "No MCP servers configured", nil)
return nil return nil
} }
logger.InfoCF("mcp", "Initializing MCP servers", logger.InfoCF(logModule, "Initializing MCP servers", map[string]any{
map[string]any{ "count": len(mcpCfg.Servers),
"count": len(mcpCfg.Servers), })
})
var wg sync.WaitGroup var wg sync.WaitGroup
errs := make(chan error, len(mcpCfg.Servers)) errs := make(chan error, len(mcpCfg.Servers))
@ -153,44 +97,29 @@ func (m *Manager) LoadFromMCPConfig(
for name, serverCfg := range mcpCfg.Servers { for name, serverCfg := range mcpCfg.Servers {
if !serverCfg.Enabled { if !serverCfg.Enabled {
logger.DebugCF("mcp", "Skipping disabled server", logger.DebugCF(logModule, "Skipping disabled server", map[string]any{
map[string]any{ "server": name,
"server": name, })
})
continue continue
} }
if err := m.validateConfig(mcpCfg); err != nil {
return fmt.Errorf("config validation failed: %w", err)
}
enabledCount++ enabledCount++
wg.Add(1) wg.Add(1)
go func(name string, serverCfg config.MCPServerConfig, workspace string) { go func(name string, serverCfg config.MCPServerConfig, workspace string) {
defer wg.Done() defer wg.Done()
// Resolve relative envFile paths relative to workspace // Resolve relative envFile paths relative to workspace
if serverCfg.EnvFile != "" && !filepath.IsAbs(serverCfg.EnvFile) { serverCfg.EnvFile = filepath.Join(workspace, serverCfg.EnvFile)
if workspace == "" {
err := fmt.Errorf(
"workspace path is empty while resolving relative envFile %q for server %s",
serverCfg.EnvFile,
name,
)
logger.ErrorCF("mcp", "Invalid MCP server configuration",
map[string]any{
"server": name,
"env_file": serverCfg.EnvFile,
"error": err.Error(),
})
errs <- err
return
}
serverCfg.EnvFile = filepath.Join(workspace, serverCfg.EnvFile)
}
if err := m.ConnectServer(ctx, name, serverCfg); err != nil { if err := m.ConnectServer(ctx, name, serverCfg); err != nil {
logger.ErrorCF("mcp", "Failed to connect to MCP server", logger.ErrorCF(logModule, "Failed to connect to MCP server", map[string]any{
map[string]any{ "server": name,
"server": name, "error": err.Error(),
"error": err.Error(), })
})
errs <- fmt.Errorf("failed to connect to server %s: %w", name, err) errs <- fmt.Errorf("failed to connect to server %s: %w", name, err)
} }
}(name, serverCfg, workspacePath) }(name, serverCfg, workspacePath)
@ -207,32 +136,51 @@ func (m *Manager) LoadFromMCPConfig(
connectedCount := len(m.GetServers()) connectedCount := len(m.GetServers())
// If all enabled servers failed to connect, return aggregated error if len(allErrors) > 0 {
if enabledCount > 0 && connectedCount == 0 { err := errors.Join(allErrors...)
logger.ErrorCF("mcp", "All MCP servers failed to connect", if connectedCount == 0 && enabledCount > 0 {
map[string]any{ logger.ErrorCF(logModule, "All MCP servers failed to connect", map[string]any{
"failed": len(allErrors), "failed": len(allErrors),
"total": enabledCount, "total": enabledCount,
}) })
return errors.Join(allErrors...) return fmt.Errorf("all MCP servers failed to connect: %w", err)
} }
if len(allErrors) > 0 { logger.WarnCF(logModule, "Initialized with partial failures", map[string]any{
logger.WarnCF("mcp", "Some MCP servers failed to connect", "failed": len(allErrors),
map[string]any{
"failed": len(allErrors),
"connected": connectedCount,
"total": enabledCount,
})
// Don't fail completely if some servers successfully connected
}
logger.InfoCF("mcp", "MCP server initialization complete",
map[string]any{
"connected": connectedCount, "connected": connectedCount,
"total": enabledCount, "total": enabledCount,
}) })
return fmt.Errorf("partial MCP initialization failure: %w", err)
}
logger.InfoCF(logModule, "MCP server initialization complete", map[string]any{
"connected": connectedCount,
"total": enabledCount,
})
return nil
}
func (m *Manager) validateConfig(mcpCfg config.MCPConfig) error {
if !mcpCfg.Enabled {
return nil
}
for name, serverCfg := range mcpCfg.Servers {
if !serverCfg.Enabled {
continue
}
if serverCfg.URL == "" && serverCfg.Command == "" {
return fmt.Errorf("server %s: missing URL (for SSE/HTTP) or command (for stdio)", name)
}
if serverCfg.EnvFile != "" && !filepath.IsAbs(serverCfg.EnvFile) {
logger.WarnCF(logModule, "Relative env file path", map[string]any{
"server_name": name,
"env_file": serverCfg.EnvFile,
})
}
}
return nil return nil
} }
@ -242,175 +190,29 @@ func (m *Manager) ConnectServer(
name string, name string,
cfg config.MCPServerConfig, cfg config.MCPServerConfig,
) error { ) error {
logger.InfoCF("mcp", "Connecting to MCP server", logger.InfoCF(logModule, "Connecting to MCP server", map[string]any{
map[string]any{ "server": name,
"server": name, "command": cfg.Command,
"command": cfg.Command, "args_count": len(cfg.Args),
"args_count": len(cfg.Args), })
})
// Create client conn, err := newServerConnection(ctx, name, cfg)
client := mcp.NewClient(&mcp.Implementation{
Name: "picoclaw",
Version: "1.0.0",
}, nil)
// Create transport based on configuration
// Auto-detect transport type if not explicitly specified
var transport mcp.Transport
transportType := cfg.Type
// Auto-detect: if URL is provided, use SSE; if command is provided, use stdio
if transportType == "" {
if cfg.URL != "" {
transportType = "sse"
} else if cfg.Command != "" {
transportType = "stdio"
} else {
return fmt.Errorf("either URL or command must be provided")
}
}
switch transportType {
case "sse", "http":
if cfg.URL == "" {
return fmt.Errorf("URL is required for SSE/HTTP transport")
}
logger.DebugCF("mcp", "Using SSE/HTTP transport",
map[string]any{
"server": name,
"url": cfg.URL,
})
sseTransport := &mcp.StreamableClientTransport{
Endpoint: cfg.URL,
}
// Add custom headers if provided
if len(cfg.Headers) > 0 {
// Create a custom HTTP client with header-injecting transport
sseTransport.HTTPClient = &http.Client{
Transport: &headerTransport{
base: http.DefaultTransport,
headers: cfg.Headers,
},
}
logger.DebugCF("mcp", "Added custom HTTP headers",
map[string]any{
"server": name,
"header_count": len(cfg.Headers),
})
}
transport = sseTransport
case "stdio":
if cfg.Command == "" {
return fmt.Errorf("command is required for stdio transport")
}
logger.DebugCF("mcp", "Using stdio transport",
map[string]any{
"server": name,
"command": cfg.Command,
})
// Create command with context
cmd := exec.CommandContext(ctx, cfg.Command, cfg.Args...)
// Build environment variables with proper override semantics
// Use a map to ensure config variables override file variables
envMap := make(map[string]string)
// Start with parent process environment
for _, e := range cmd.Environ() {
if idx := strings.Index(e, "="); idx > 0 {
envMap[e[:idx]] = e[idx+1:]
}
}
// Load environment variables from file if specified
if cfg.EnvFile != "" {
envVars, err := loadEnvFile(cfg.EnvFile)
if err != nil {
return fmt.Errorf("failed to load env file %s: %w", cfg.EnvFile, err)
}
for k, v := range envVars {
envMap[k] = v
}
logger.DebugCF("mcp", "Loaded environment variables from file",
map[string]any{
"server": name,
"envFile": cfg.EnvFile,
"var_count": len(envVars),
})
}
// Environment variables from config override those from file
for k, v := range cfg.Env {
envMap[k] = v
}
// Convert map to slice
env := make([]string, 0, len(envMap))
for k, v := range envMap {
env = append(env, fmt.Sprintf("%s=%s", k, v))
}
cmd.Env = env
transport = &mcp.CommandTransport{Command: cmd}
default:
return fmt.Errorf(
"unsupported transport type: %s (supported: stdio, sse, http)",
transportType,
)
}
// Connect to server
session, err := client.Connect(ctx, transport, nil)
if err != nil { if err != nil {
return fmt.Errorf("failed to connect: %w", err) return fmt.Errorf("failed to create server connection: %w", err)
} }
// Get server info
initResult := session.InitializeResult()
logger.InfoCF("mcp", "Connected to MCP server",
map[string]any{
"server": name,
"serverName": initResult.ServerInfo.Name,
"serverVersion": initResult.ServerInfo.Version,
"protocol": initResult.ProtocolVersion,
})
// List available tools if supported
var tools []*mcp.Tool
if initResult.Capabilities.Tools != nil {
for tool, err := range session.Tools(ctx, nil) {
if err != nil {
logger.WarnCF("mcp", "Error listing tool",
map[string]any{
"server": name,
"error": err.Error(),
})
continue
}
tools = append(tools, tool)
}
logger.InfoCF("mcp", "Listed tools from MCP server",
map[string]any{
"server": name,
"toolCount": len(tools),
})
}
// Store connection
m.mu.Lock() m.mu.Lock()
m.servers[name] = &ServerConnection{ if oldConn, exists := m.servers[name]; exists {
Name: name, logger.WarnCF(logModule, "Overwriting existing server connection, closing old session",
Client: client, map[string]any{"server": name},
Session: session, )
Tools: tools, _ = oldConn.Session.Close()
} }
m.servers[name] = conn
m.mu.Unlock() m.mu.Unlock()
// start health monitoring for this server
m.startMonitor(name)
return nil return nil
} }
@ -420,9 +222,7 @@ func (m *Manager) GetServers() map[string]*ServerConnection {
defer m.mu.RUnlock() defer m.mu.RUnlock()
result := make(map[string]*ServerConnection, len(m.servers)) result := make(map[string]*ServerConnection, len(m.servers))
for k, v := range m.servers { maps.Copy(result, m.servers)
result[k] = v
}
return result return result
} }
@ -441,26 +241,20 @@ func (m *Manager) CallTool(
serverName, toolName string, serverName, toolName string,
arguments map[string]any, arguments map[string]any,
) (*mcp.CallToolResult, error) { ) (*mcp.CallToolResult, error) {
// Check if closed before acquiring lock (fast path)
if m.closed.Load() {
return nil, fmt.Errorf("manager is closed")
}
m.mu.RLock() m.mu.RLock()
// Double-check after acquiring lock to prevent TOCTOU race
if m.closed.Load() { if m.closed.Load() {
m.mu.RUnlock() m.mu.RUnlock()
return nil, fmt.Errorf("manager is closed") return nil, fmt.Errorf("manager is closed")
} }
conn, ok := m.servers[serverName]
if ok {
m.wg.Add(1) // Add to WaitGroup while holding the lock
}
m.mu.RUnlock()
conn, ok := m.servers[serverName]
if !ok { if !ok {
m.mu.RUnlock()
return nil, fmt.Errorf("server %s not found", serverName) return nil, fmt.Errorf("server %s not found", serverName)
} }
m.wg.Add(1)
m.mu.RUnlock()
defer m.wg.Done() defer m.wg.Done()
params := &mcp.CallToolParams{ params := &mcp.CallToolParams{
@ -468,43 +262,42 @@ func (m *Manager) CallTool(
Arguments: arguments, Arguments: arguments,
} }
result, err := conn.Session.CallTool(ctx, params) return conn.Session.CallTool(ctx, params)
if err != nil {
return nil, fmt.Errorf("failed to call tool: %w", err)
}
return result, nil
} }
// Close closes all server connections // Close closes all server connections
func (m *Manager) Close() error { func (m *Manager) Close() error {
// Use Swap to atomically set closed=true and get the previous value m.mu.Lock()
// This prevents TOCTOU race with CallTool's closed check
if m.closed.Swap(true) { if m.closed.Swap(true) {
m.mu.Unlock()
return nil // already closed return nil // already closed
} }
// Wait for all in-flight CallTool calls to finish before closing sessions m.mu.Unlock()
// After closed=true is set, no new CallTool can start (they check closed first)
m.wg.Wait() m.wg.Wait()
m.mu.Lock() m.mu.Lock()
defer m.mu.Unlock() defer m.mu.Unlock()
logger.InfoCF("mcp", "Closing all MCP server connections", logger.InfoCF(logModule, "Closing all MCP server connections", map[string]any{
map[string]any{ "count": len(m.servers),
"count": len(m.servers), })
})
var errs []error var errs []error
for name, conn := range m.servers { for name, conn := range m.servers {
if err := conn.Session.Close(); err != nil { if conn.cancelFunc != nil {
logger.ErrorCF("mcp", "Failed to close server connection", conn.cancelFunc()
map[string]any{ }
if conn.Session != nil {
if err := conn.Session.Close(); err != nil {
logger.ErrorCF(logModule, "Failed to close server connection", map[string]any{
"server": name, "server": name,
"error": err.Error(), "error": err.Error(),
}) })
errs = append(errs, fmt.Errorf("server %s: %w", name, err)) errs = append(errs, fmt.Errorf("server %s: %w", name, err))
}
} }
} }
@ -530,3 +323,100 @@ func (m *Manager) GetAllTools() map[string][]*mcp.Tool {
} }
return result return result
} }
func (m *Manager) startMonitor(name string) {
ctx, cancel := context.WithCancel(context.Background())
m.mu.Lock()
if conn, ok := m.servers[name]; ok {
if conn.cancelFunc != nil {
conn.cancelFunc()
}
conn.cancelFunc = cancel
}
m.mu.Unlock()
go func() {
defer cancel()
ticker := time.NewTicker(30 * time.Second)
defer ticker.Stop()
count := 0
for {
select {
case <-ctx.Done():
return
case <-ticker.C:
if !m.checkHealth(name) {
count++
if count >= maxFails {
m.handleServerOffline(name)
return
}
} else {
count = 0
}
}
}
}()
}
// checkHealth performs a health check by calling ListTools. If it fails, it returns false.
func (m *Manager) checkHealth(name string) bool {
conn, ok := m.GetServer(name)
if !ok || conn.Session == nil {
return false
}
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
_, err := conn.Session.ListTools(ctx, nil)
return err == nil
}
// handleServerOffline marks the server as offline and starts async reconnection attempts
func (m *Manager) handleServerOffline(name string) {
m.mu.RLock()
conn, ok := m.servers[name]
m.mu.RUnlock()
if !ok {
return
}
conn.status.Store(StatusOffline)
logger.WarnCF("mcp", "Server is offline, starting async reconnection", map[string]any{"server": name})
go func() {
backoff := []time.Duration{
1 * time.Second,
2 * time.Second,
4 * time.Second,
8 * time.Second,
16 * time.Second,
}
attempt := 0
for {
wait := backoff[len(backoff)-1]
if attempt < len(backoff) {
wait = backoff[attempt]
}
time.Sleep(wait)
attempt++
logger.DebugCF("mcp", "Reconnection attempt", map[string]any{"server": name, "attempt": attempt})
err := m.ConnectServer(context.Background(), name, conn.Config)
if err == nil {
logger.InfoCF("mcp", "Reconnection successful", map[string]any{"server": name})
conn.status.Store(StatusOnline)
m.startMonitor(name)
return
}
}
}()
}

View file

@ -2,11 +2,13 @@ package mcp
import ( import (
"context" "context"
"maps"
"os" "os"
"path/filepath" "path/filepath"
"strings" "strings"
"testing" "testing"
"github.com/joho/godotenv"
sdkmcp "github.com/modelcontextprotocol/go-sdk/mcp" sdkmcp "github.com/modelcontextprotocol/go-sdk/mcp"
"github.com/sipeed/picoclaw/pkg/config" "github.com/sipeed/picoclaw/pkg/config"
@ -100,10 +102,10 @@ PORT =8080`,
t.Fatalf("Failed to create test file: %v", err) t.Fatalf("Failed to create test file: %v", err)
} }
result, err := loadEnvFile(envFile) result, err := godotenv.Read(envFile)
if tt.expectErr { if tt.expectErr {
if err == nil { if err != nil && result[""] == tt.content {
t.Errorf("Expected error but got none") t.Errorf("Expected error but got none")
} }
return return
@ -130,7 +132,7 @@ PORT =8080`,
} }
func TestLoadEnvFileNotFound(t *testing.T) { func TestLoadEnvFileNotFound(t *testing.T) {
_, err := loadEnvFile("/nonexistent/file.env") _, err := godotenv.Read("/nonexistent/file.env")
if err == nil { if err == nil {
t.Error("Expected error for nonexistent file") t.Error("Expected error for nonexistent file")
} }
@ -150,7 +152,7 @@ SHARED_VAR=from_file`
} }
// Load envFile // Load envFile
envVars, err := loadEnvFile(envFile) envVars, err := godotenv.Read(envFile)
if err != nil { if err != nil {
t.Fatalf("Failed to load env file: %v", err) t.Fatalf("Failed to load env file: %v", err)
} }
@ -168,12 +170,8 @@ SHARED_VAR=from_file`
// Merge: envFile first, then config overrides // Merge: envFile first, then config overrides
merged := make(map[string]string) merged := make(map[string]string)
for k, v := range envVars { maps.Copy(merged, envVars)
merged[k] = v maps.Copy(merged, configEnv)
}
for k, v := range configEnv {
merged[k] = v
}
// Verify priority: config.Env should override envFile // Verify priority: config.Env should override envFile
if merged["SHARED_VAR"] != "from_config" { if merged["SHARED_VAR"] != "from_config" {
@ -212,8 +210,8 @@ func TestLoadFromMCPConfig_EmptyWorkspaceWithRelativeEnvFile(t *testing.T) {
t.Fatal("expected error for relative env_file with empty workspace path, got nil") t.Fatal("expected error for relative env_file with empty workspace path, got nil")
} }
if !strings.Contains(err.Error(), "workspace path is empty") { if !strings.Contains(err.Error(), "failed to load env file") {
t.Fatalf("expected workspace path validation error, got: %v", err) t.Fatalf("failed to load env file, got: %v", err)
} }
} }

236
pkg/mcp/server.go Normal file
View file

@ -0,0 +1,236 @@
package mcp
import (
"context"
"fmt"
"maps"
"net/http"
"os/exec"
"strings"
"sync"
"sync/atomic"
"syscall"
"github.com/joho/godotenv"
"github.com/modelcontextprotocol/go-sdk/mcp"
"github.com/sipeed/picoclaw/pkg/config"
"github.com/sipeed/picoclaw/pkg/logger"
)
const (
TransportTypeSSE = "sse"
TransportTypeHTTP = "http"
TransportTypeStdio = "stdio"
)
type ServerStatus int
const (
StatusOnline ServerStatus = iota
StatusOffline
StatusConnecting
)
// ServerConnection represents a connection to an MCP server
type ServerConnection struct {
Name string
Config config.MCPServerConfig // save config for potential reconnection
Client *mcp.Client
Session *mcp.ClientSession
Tools []*mcp.Tool
status atomic.Value // ServerStatus
mu sync.Mutex // protect session switching
cancelFunc context.CancelFunc // for canceling ongoing operations during reconnection
}
func newServerConnection(
ctx context.Context,
name string,
cfg config.MCPServerConfig,
) (*ServerConnection, error) {
conn := &ServerConnection{
Name: name,
Config: cfg,
Client: mcp.NewClient(&mcp.Implementation{
Name: clientName,
Version: clientVersion,
}, nil),
mu: sync.Mutex{},
}
transport, err := conn.createTransport(cfg)
if err != nil {
return nil, fmt.Errorf("failed to create transport: %w", err)
}
session, err := conn.Client.Connect(ctx, transport, nil)
if err != nil {
return nil, fmt.Errorf("failed to connect: %w", err)
}
initResult := session.InitializeResult()
logger.InfoCF(logModule, "Connected to MCP server", map[string]any{
"server": name,
"serverName": initResult.ServerInfo.Name,
"serverVersion": initResult.ServerInfo.Version,
"protocol": initResult.ProtocolVersion,
})
var tools []*mcp.Tool
if initResult.Capabilities.Tools != nil {
for tool, err := range session.Tools(ctx, nil) {
if err != nil {
logger.WarnCF(logModule, "Error listing tool", map[string]any{
"server": name,
"error": err.Error(),
})
continue
}
tools = append(tools, tool)
}
logger.InfoCF(logModule, "Listed tools from MCP server", map[string]any{
"server": name,
"toolCount": len(tools),
})
}
conn.Tools = tools
conn.Session = session
return conn, nil
}
func (conn *ServerConnection) createTransport(cfg config.MCPServerConfig) (mcp.Transport, error) {
transportType := conn.detectTransportType(cfg)
switch transportType {
case TransportTypeSSE, TransportTypeHTTP:
return conn.newSSETransport(context.Background(), "temp", cfg)
case TransportTypeStdio:
return conn.newStdioTransport(context.Background(), "temp", cfg)
default:
return nil, fmt.Errorf("unsupported transport type: %s (supported: stdio, sse, http)", transportType)
}
}
func (conn *ServerConnection) detectTransportType(cfg config.MCPServerConfig) string {
if cfg.Type != "" {
return cfg.Type
}
if cfg.URL != "" {
return TransportTypeSSE
}
if cfg.Command != "" {
return TransportTypeStdio
}
return ""
}
// Build StdioTransport
func (conn *ServerConnection) newStdioTransport(
ctx context.Context,
name string,
cfg config.MCPServerConfig,
) (mcp.Transport, error) {
select {
case <-ctx.Done():
return nil, ctx.Err()
default:
if cfg.Command == "" {
return nil, ErrStdioCommandRequired
}
logger.DebugCF(logModule, "Using stdio transport", map[string]any{
"server": name,
"command": cfg.Command,
"args": cfg.Args,
})
cmd := exec.CommandContext(ctx, cfg.Command, cfg.Args...)
cmd.SysProcAttr = &syscall.SysProcAttr{
Setpgid: true,
}
cmdEnv := cmd.Environ()
envMap := make(map[string]string, len(cmdEnv)/2)
for _, e := range cmdEnv {
if idx := strings.SplitN(e, "=", 2); len(idx) == 2 {
envMap[idx[0]] = idx[1]
}
}
if cfg.EnvFile != "" {
envVars, err := godotenv.Read(cfg.EnvFile)
if err != nil {
return nil, fmt.Errorf("failed to load env file %s: %w", cfg.EnvFile, err)
}
maps.Copy(envMap, envVars)
logger.DebugCF(logModule, "Loaded environment variables from file", map[string]any{
"server": name,
"envFile": cfg.EnvFile,
"args": envVars,
})
}
maps.Copy(envMap, cfg.Env)
env := make([]string, 0, len(envMap))
for k, v := range envMap {
env = append(env, k+"="+v)
}
cmd.Env = make([]string, len(env))
copy(cmd.Env, env)
transport := &mcp.CommandTransport{Command: cmd}
return transport, nil
}
}
// Build SSETransport
func (conn *ServerConnection) newSSETransport(
ctx context.Context,
name string,
cfg config.MCPServerConfig,
) (mcp.Transport, error) {
select {
case <-ctx.Done():
return nil, ctx.Err()
default:
if cfg.URL == "" {
return nil, ErrInvalidServerConfig
}
logger.DebugCF(logModule, "Using SSE/HTTP transport", map[string]any{
"server": name,
"url": cfg.URL,
})
sseTransport := &mcp.StreamableClientTransport{
Endpoint: cfg.URL,
}
if len(cfg.Headers) > 0 {
sseTransport.HTTPClient = &http.Client{
Transport: &headerTransport{
base: http.DefaultTransport,
headers: cfg.Headers,
},
}
logger.DebugCF(logModule, "Added custom HTTP headers", map[string]any{
"server": name,
"header_count": len(cfg.Headers),
})
}
return sseTransport, nil
}
}
func (conn *ServerConnection) GetStatus() ServerStatus {
return conn.status.Load().(ServerStatus)
}