fix(mcp): return aggregated error when all servers fail to connect

- Add errors.Join to return aggregated error when all enabled MCP servers fail
- Track enabled server count separately from total configured servers
- Return error only when all servers fail, not for partial failures
- Improve logging with accurate server counts (enabled vs connected)
- Maintains fault tolerance: partial failures don't stop initialization
This commit is contained in:
yuchou87 2026-02-16 19:33:31 +08:00
parent 27f5494038
commit abd4830e80

View file

@ -3,6 +3,7 @@ package mcp
import ( import (
"bufio" "bufio"
"context" "context"
"errors"
"fmt" "fmt"
"net/http" "net/http"
"os" "os"
@ -24,12 +25,12 @@ 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 // Clone the request to avoid modifying the original
req = req.Clone(req.Context()) req = req.Clone(req.Context())
// Add custom headers // Add custom headers
for key, value := range t.headers { for key, value := range t.headers {
req.Header.Set(key, value) req.Header.Set(key, value)
} }
// Use the base transport // Use the base transport
base := t.base base := t.base
if base == nil { if base == nil {
@ -129,6 +130,7 @@ func (m *Manager) LoadFromConfig(ctx context.Context, cfg *config.Config) error
var wg sync.WaitGroup var wg sync.WaitGroup
errs := make(chan error, len(cfg.Tools.MCP.Servers)) errs := make(chan error, len(cfg.Tools.MCP.Servers))
enabledCount := 0
for name, serverCfg := range cfg.Tools.MCP.Servers { for name, serverCfg := range cfg.Tools.MCP.Servers {
if !serverCfg.Enabled { if !serverCfg.Enabled {
@ -139,6 +141,7 @@ func (m *Manager) LoadFromConfig(ctx context.Context, cfg *config.Config) error
continue continue
} }
enabledCount++
wg.Add(1) wg.Add(1)
go func(name string, serverCfg config.MCPServerConfig) { go func(name string, serverCfg config.MCPServerConfig) {
defer wg.Done() defer wg.Done()
@ -163,20 +166,32 @@ func (m *Manager) LoadFromConfig(ctx context.Context, cfg *config.Config) error
allErrors = append(allErrors, err) allErrors = append(allErrors, err)
} }
connectedCount := len(m.GetServers())
// If all enabled servers failed to connect, return aggregated error
if enabledCount > 0 && connectedCount == 0 {
logger.ErrorCF("mcp", "All MCP servers failed to connect",
map[string]interface{}{
"failed": len(allErrors),
"total": enabledCount,
})
return errors.Join(allErrors...)
}
if len(allErrors) > 0 { if len(allErrors) > 0 {
logger.WarnCF("mcp", "Some MCP servers failed to connect", logger.WarnCF("mcp", "Some MCP servers failed to connect",
map[string]interface{}{ map[string]interface{}{
"failed": len(allErrors), "failed": len(allErrors),
"total": len(cfg.Tools.MCP.Servers), "connected": connectedCount,
"total": enabledCount,
}) })
// Don't fail completely if some servers fail to connect // Don't fail completely if some servers successfully connected
} }
connectedCount := len(m.GetServers())
logger.InfoCF("mcp", "MCP server initialization complete", logger.InfoCF("mcp", "MCP server initialization complete",
map[string]interface{}{ map[string]interface{}{
"connected": connectedCount, "connected": connectedCount,
"total": len(cfg.Tools.MCP.Servers), "total": enabledCount,
}) })
return nil return nil
@ -223,11 +238,11 @@ func (m *Manager) ConnectServer(ctx context.Context, name string, cfg config.MCP
"server": name, "server": name,
"url": cfg.URL, "url": cfg.URL,
}) })
sseTransport := &mcp.StreamableClientTransport{ sseTransport := &mcp.StreamableClientTransport{
Endpoint: cfg.URL, Endpoint: cfg.URL,
} }
// Add custom headers if provided // Add custom headers if provided
if len(cfg.Headers) > 0 { if len(cfg.Headers) > 0 {
// Create a custom HTTP client with header-injecting transport // Create a custom HTTP client with header-injecting transport
@ -243,7 +258,7 @@ func (m *Manager) ConnectServer(ctx context.Context, name string, cfg config.MCP
"header_count": len(cfg.Headers), "header_count": len(cfg.Headers),
}) })
} }
transport = sseTransport transport = sseTransport
case "stdio": case "stdio":
if cfg.Command == "" { if cfg.Command == "" {
@ -259,7 +274,7 @@ func (m *Manager) ConnectServer(ctx context.Context, name string, cfg config.MCP
// Set environment variables // Set environment variables
env := cmd.Environ() env := cmd.Environ()
// Load environment variables from file if specified // Load environment variables from file if specified
if cfg.EnvFile != "" { if cfg.EnvFile != "" {
envVars, err := loadEnvFile(cfg.EnvFile) envVars, err := loadEnvFile(cfg.EnvFile)
@ -276,14 +291,14 @@ func (m *Manager) ConnectServer(ctx context.Context, name string, cfg config.MCP
"var_count": len(envVars), "var_count": len(envVars),
}) })
} }
// Environment variables from config override those from file // Environment variables from config override those from file
if len(cfg.Env) > 0 { if len(cfg.Env) > 0 {
for k, v := range cfg.Env { for k, v := range cfg.Env {
env = append(env, fmt.Sprintf("%s=%s", k, v)) env = append(env, fmt.Sprintf("%s=%s", k, v))
} }
} }
// Set environment if we added any variables // Set environment if we added any variables
if len(env) > len(cmd.Environ()) { if len(env) > len(cmd.Environ()) {
cmd.Env = env cmd.Env = env