Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
149 changes: 139 additions & 10 deletions internal/upstream/managed/client.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,8 +4,10 @@ import (
"context"
"errors"
"fmt"
"io"
"net"
"net/url"
"os"
"strings"
"sync"
"sync/atomic"
Expand All @@ -22,6 +24,7 @@ import (
"github.com/smart-mcp-proxy/mcpproxy-go/internal/upstream/core"
"github.com/smart-mcp-proxy/mcpproxy-go/internal/upstream/types"

mcptransport "github.com/mark3labs/mcp-go/client/transport"
"github.com/mark3labs/mcp-go/mcp"
"go.uber.org/zap"
)
Expand Down Expand Up @@ -1058,6 +1061,7 @@ func (mc *Client) runListToolsAsLeader(listCtx context.Context, release func() b
}
}()

listEpoch := mc.connectionEpoch.Load()
tools, err := mc.coreClient.ListTools(listCtx)
mc.publishListToolsResult(tools, err)

Expand All @@ -1074,7 +1078,9 @@ func (mc *Client) runListToolsAsLeader(listCtx context.Context, release func() b
mc.logger.Warn("Connection error detected during ListTools, updating server state",
zap.String("server", mc.GetConfig().Name),
zap.Error(err))
mc.StateManager.SetError(err)
// Guarded like the call and health paths: only the connection
// that produced the failure may be marked, once.
mc.setErrorIfCurrentConnection(err, listEpoch)
}
return nil, fmt.Errorf("ListTools failed: %w", err)
}
Expand Down Expand Up @@ -1138,7 +1144,6 @@ func (mc *Client) callTool(ctx context.Context, toolName string, args map[string
if !mc.IsConnected() {
return nil, fmt.Errorf("client not connected (state: %s)", mc.StateManager.GetState().String())
}

// #1317 round 2: counted from HERE, before admission control, not just
// around the transport call below — a call already queued in
// acquireAdmission's wait is just as "in flight" from tryReconnect()'s
Expand Down Expand Up @@ -1192,6 +1197,10 @@ func (mc *Client) callTool(ctx context.Context, toolName string, args map[string
invoker = mc.coreClient
}

// The connection generation this call runs on, read AFTER the admission
// wait (a queued call may outlive a reconnect), so a transport failure is
// only ever charged to the session that produced it (RC4-STDIO-001).
callEpoch := mc.connectionEpoch.Load()
result, err := invoker.CallTool(ctx, toolName, args)
if err != nil {
mc.recordCallToolOAuthSignal(toolName, err)
Expand Down Expand Up @@ -1247,7 +1256,7 @@ func (mc *Client) callTool(ctx context.Context, toolName string, args map[string
zap.String("tool", toolName),
zap.Error(err))
}
mc.StateManager.SetError(err)
mc.recordDeadConnection(toolName, err, callEpoch)
}
return nil, err
}
Expand All @@ -1262,6 +1271,47 @@ func (mc *Client) callTool(ctx context.Context, toolName string, args map[string
return result, nil
}

// recordDeadConnection flips the server to Error after a tools/call produced
// hard evidence that its connection is broken (refused/reset, broken pipe, or
// a dead stdio transport). The health loop's next tick then reconnects through
// the normal ShouldRetry backoff.
//
// The verdict is applied only while the client is still Ready on the SAME
// connection generation the call ran on, serialized with Connect/Disconnect
// through epochMu (the same pairing probeAfterAmbiguousCallError uses):
// - a call that was on the wire when the user disconnected the server sees
// "transport closed" too; that belongs to the closed generation and must
// not flip the Disconnected client to Error;
// - a burst of concurrent calls hitting the same dead transport records ONE
// failure, not one per call — every extra SetError would bump retryCount
// and stretch the reconnect backoff for a single process death.
func (mc *Client) recordDeadConnection(toolName string, err error, callEpoch int64) {
if !mc.setErrorIfCurrentConnection(err, callEpoch) {
mc.logger.Debug("Tool call connection failure belongs to a connection that is already gone; not re-marking",
zap.String("server", mc.GetConfig().Name),
zap.String("tool", toolName),
zap.Error(err))
}
}

// setErrorIfCurrentConnection applies SetError(err) only while the client is
// still Ready on connection generation epoch, serialized with Connect and
// Disconnect through epochMu. It reports whether the error was recorded.
//
// Both the tools/call path and the health loop use it: Disconnect closes the
// transport (every in-flight request then fails with "transport closed")
// BEFORE it bumps the epoch and Resets the state, so a verdict computed from
// that failure must not land on the freshly Disconnected client.
func (mc *Client) setErrorIfCurrentConnection(err error, epoch int64) bool {
mc.epochMu.Lock()
defer mc.epochMu.Unlock()
if !mc.IsConnected() || mc.connectionEpoch.Load() != epoch {
return false
}
mc.StateManager.SetError(err)
return true
}

// recordCallToolOAuthSignal inspects a failed tools/call error and, when it is an
// "authorization required" / 401 from an otherwise-connected server (NOT a
// connection error), flags the server as needing OAuth sign-in. This covers
Expand Down Expand Up @@ -1698,6 +1748,10 @@ func (mc *Client) performHealthCheck() {
if prober == nil {
prober = mc.coreClient
}
// The generation this ping runs on: a Disconnect racing the ping closes
// the transport (the ping fails with "transport closed") and only then
// bumps the epoch, so the verdict below must land on this generation only.
pingEpoch := mc.connectionEpoch.Load()
err := prober.Ping(ctx)

if err != nil {
Expand Down Expand Up @@ -1747,11 +1801,16 @@ func (mc *Client) performHealthCheck() {
mc.resetInFlightSuppression()
}
if mc.recordHealthCheckFailure(err) {
mc.logger.Warn("Health check failed repeatedly, marking as error",
zap.String("server", mc.GetConfig().Name),
zap.Int("consecutive_failures", mc.consecutiveHealthFailures),
zap.Error(err))
mc.StateManager.SetError(err)
if mc.setErrorIfCurrentConnection(err, pingEpoch) {
mc.logger.Warn("Health check failed repeatedly, marking as error",
zap.String("server", mc.GetConfig().Name),
zap.Int("consecutive_failures", mc.consecutiveHealthFailures),
zap.Error(err))
} else {
mc.logger.Debug("Health check failure belongs to a connection that is already gone; not marking",
zap.String("server", mc.GetConfig().Name),
zap.Error(err))
}
} else {
mc.logger.Info("Health check failed transiently, tolerating below threshold",
zap.String("server", mc.GetConfig().Name),
Expand Down Expand Up @@ -1836,6 +1895,12 @@ func isTransientHealthCheckError(err error) bool {
if err == nil {
return false
}
// A dead transport (the stdio child exited, its pipes closed) never
// recovers by waiting: every later request fails the same way until the
// process is respawned (RC4-STDIO-001).
if isDeadTransportError(err) {
return false
}
msg := strings.ToLower(err.Error())
// Hard failures: short-circuit to "not transient" so the caller flips
// Error on the first miss. Order matters — check these BEFORE the
Expand Down Expand Up @@ -2133,11 +2198,73 @@ func (mc *Client) waitForReconnectCompletion(ctx context.Context) error {
}
}

// isDeadTransportError reports whether err proves the transport itself is
// closed, as opposed to one request failing on a live connection. The stdio
// case is RC4-STDIO-001: when the upstream's child process dies, mcp-go's
// stdio reader hits EOF and closes the transport, and every later request
// returns transport.ErrTransportClosed ("transport closed") — which matched
// none of isConnectionError's markers, so the health loop ignored it as a
// "high activity timeout" and the server stayed Ready with every call failing.
//
// The classifier is shared by every protocol, so it only accepts evidence that
// the LOCAL end of the transport is closed:
// - io.EOF / io.ErrUnexpectedEOF are deliberately NOT matched. net/http wraps
// them for a single truncated or reset response on a live HTTP/SSE
// upstream (LB idle timeout, keep-alive race), which must stay subject to
// the normal flap tolerance. A stdio child's death never surfaces as them:
// mcp-go's stdio reader turns its EOF into ErrTransportClosed.
// - The text fallbacks (for chains flattened before they reach us) only
// count after mcp-go's "transport error: " prefix, i.e. text produced by
// transport.Error — never a live server's JSON-RPC error message, which
// mcp-go returns unwrapped and which may echo errno text such as
// "file already closed".
func isDeadTransportError(err error) bool {
if err == nil {
return false
}
if errors.Is(err, mcptransport.ErrTransportClosed) ||
errors.Is(err, io.ErrClosedPipe) ||
errors.Is(err, os.ErrClosed) {
return true
}
// Text fallback, for chains that were flattened into a string: only the
// payload of mcp-go's own transport wrapper counts — either the typed
// *transport.Error, or a message that STARTS with its prefix. A server's
// JSON-RPC error is surfaced unwrapped (errors.New(message)), so text it
// supplies can only match if it is the entire message, not by containing
// the phrase somewhere.
var transportMsg string
var tErr *mcptransport.Error
if errors.As(err, &tErr) && tErr.Err != nil {
transportMsg = tErr.Err.Error()
} else {
const transportPrefix = "transport error: "
msg := err.Error()
if !strings.HasPrefix(msg, transportPrefix) {
return false
}
transportMsg = strings.TrimPrefix(msg, transportPrefix)
}
for _, marker := range []string{
mcptransport.ErrTransportClosed.Error(), // "transport closed"
os.ErrClosed.Error(), // "file already closed"
io.ErrClosedPipe.Error(), // "io: read/write on closed pipe"
} {
if strings.Contains(transportMsg, marker) {
return true
}
}
return false
}

// isConnectionError checks if an error indicates a connection problem
func (mc *Client) isConnectionError(err error) bool {
if err == nil {
return false
}
if isDeadTransportError(err) {
return true
}

errStr := err.Error()
connectionErrors := []string{
Expand Down Expand Up @@ -2352,6 +2479,7 @@ func (mc *Client) GetCachedToolCount(ctx context.Context) (int, error) {

// Fetch fresh tool count with timeout. Publish the result so any concurrent
// ListTools waiter coalesced behind us receives the real tools list.
countEpoch := mc.connectionEpoch.Load()
tools, err := mc.coreClient.ListTools(listCtx)
mc.publishListToolsResult(tools, err)
if err != nil {
Expand All @@ -2360,9 +2488,10 @@ func (mc *Client) GetCachedToolCount(ctx context.Context) (int, error) {
zap.Error(err),
zap.Int("cached_count", cachedCount))

// Check if it's a connection error and update state
// Check if it's a connection error and update state (guarded: only
// the connection that produced the failure may be marked, once).
if mc.isConnectionError(err) {
mc.StateManager.SetError(err)
mc.setErrorIfCurrentConnection(err, countEpoch)
}

// Return cached count if available, even if stale
Expand Down
110 changes: 110 additions & 0 deletions internal/upstream/managed/dead_transport_fixture_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,110 @@
package managed

import (
"context"
"encoding/json"
"os/exec"
"path/filepath"
"runtime"
"testing"
"time"

"github.com/mark3labs/mcp-go/mcp"
"github.com/stretchr/testify/require"
"go.uber.org/zap"

"github.com/smart-mcp-proxy/mcpproxy-go/internal/config"
"github.com/smart-mcp-proxy/mcpproxy-go/internal/secret"
)

// pingInstanceID calls the fixture's deterministic `ping` tool and returns the
// per-process instance_id, which changes on every process start.
func pingInstanceID(ctx context.Context, mc *Client) (string, error) {
res, err := mc.CallTool(ctx, "ping", nil)
if err != nil {
return "", err
}
for _, c := range res.Content {
if tc, ok := c.(mcp.TextContent); ok {
var body struct {
InstanceID string `json:"instance_id"`
}
if jerr := json.Unmarshal([]byte(tc.Text), &body); jerr == nil && body.InstanceID != "" {
return body.InstanceID, nil
}
}
}
return "", nil
}

// TestStdioChildKilled_ClientLeavesReadyAndRespawns is the end-to-end
// RC4-STDIO-001 regression: an admitted stdio upstream (cmd/mcpfixture) whose
// child process is SIGKILLed must stop reporting Ready on the first failed
// call ("transport closed") and be respawned by the managed client's own
// health loop + backoff, without any explicit restart.
func TestStdioChildKilled_ClientLeavesReadyAndRespawns(t *testing.T) {
if testing.Short() {
t.Skip("builds and spawns the mcpfixture stdio binary")
}
if runtime.GOOS == "windows" {
t.Skip("uses pkill to SIGKILL the child")
}
if _, err := exec.LookPath("pkill"); err != nil {
t.Skip("pkill not available")
}

bin := filepath.Join(t.TempDir(), "mcpfixture-deadtransport")
build := exec.Command("go", "build", "-o", bin, "github.com/smart-mcp-proxy/mcpproxy-go/cmd/mcpfixture")
out, err := build.CombinedOutput()
require.NoError(t, err, "build mcpfixture: %s", out)

interval := config.Duration(200 * time.Millisecond)
cfg := &config.ServerConfig{
Name: "stdio-dead-transport",
Command: bin,
Args: []string{"--transport", "stdio"},
Protocol: "stdio",
Enabled: true,
HealthCheckInterval: &interval,
}
mc, err := NewClient(cfg.Name, cfg, zap.NewNop(), nil, &config.Config{}, nil, secret.NewResolver())
require.NoError(t, err)

ctx, cancel := context.WithTimeout(context.Background(), 60*time.Second)
defer cancel()
require.NoError(t, mc.Connect(ctx))
t.Cleanup(func() { _ = mc.Disconnect() })

firstID, err := pingInstanceID(ctx, mc)
require.NoError(t, err)
require.NotEmpty(t, firstID)
epochBefore := mc.ConnectionEpoch()

// SIGKILL only the child: the proxy-side client is left untouched.
require.NoError(t, exec.Command("pkill", "-KILL", "-f", bin).Run())

// The dead transport must surface as a failed call, and that call must
// stop the client from reporting Ready on the dead generation.
require.Eventually(t, func() bool {
_, callErr := mc.CallTool(ctx, "ping", nil)
return callErr != nil
}, 5*time.Second, 20*time.Millisecond, "calls must start failing once the child is dead")
require.True(t, !mc.IsConnected() || mc.ConnectionEpoch() != epochBefore,
"after a call failed on the dead transport the client must not still report Ready on that connection (state=%s)",
mc.GetState())

// ...and the managed client must respawn the process on its own.
var newID string
require.Eventually(t, func() bool {
if !mc.IsConnected() {
return false
}
id, callErr := pingInstanceID(ctx, mc)
if callErr != nil || id == "" {
return false
}
newID = id
return true
}, 20*time.Second, 100*time.Millisecond, "the client must reconnect a fresh stdio process (state=%s)", mc.GetState())
require.NotEqual(t, firstID, newID, "the reconnected upstream must be a new process")
}
Loading
Loading