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:
parent
27f5494038
commit
abd4830e80
1 changed files with 28 additions and 13 deletions
|
|
@ -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
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue