diff --git a/docs/setup.md b/docs/setup.md index baf1de4a4..6b69e9fea 100644 --- a/docs/setup.md +++ b/docs/setup.md @@ -757,6 +757,10 @@ Get-Content -Path "$env:LOCALAPPDATA\mcpproxy\logs\main.log" -Wait tail -f ~/Library/Logs/mcpproxy/main.log | grep -E "(github-server|oauth|error)" ``` +**Upstream session lost (HTTP 404):** + +A remote Streamable HTTP server that restarts or expires its session answers HTTP 404 to the session id MCPProxy holds. MCPProxy re-initializes the MCP session on the same connection instead of marking the server as errored. List, prompt and health-ping requests are retried once. A `tools/call` is retried once only for read-only tools (`readOnlyHint: true`) whose description and schema did not change; any other tool returns an error saying the session was re-established and the call was not repeated, and the next call uses the new session. Each re-initialization is logged at info level (`Upstream session terminated (HTTP 404); re-initialized in place`) and counted as `session_reinit_count` in the server's connection status. If the re-initialization or the retry also fails, the normal reconnect path applies. + ## Advanced Configuration **📚 For complete configuration reference:** See [Configuration Documentation](configuration.md) for all available options. diff --git a/internal/upstream/core/session_reinit.go b/internal/upstream/core/session_reinit.go new file mode 100644 index 000000000..e3110e60d --- /dev/null +++ b/internal/upstream/core/session_reinit.go @@ -0,0 +1,70 @@ +package core + +import ( + "context" + "fmt" + + "github.com/mark3labs/mcp-go/client/transport" + "github.com/mark3labs/mcp-go/mcp" + "go.uber.org/zap" +) + +// Spec 113-e (G8): helpers that let the managed layer re-initialize a +// Streamable HTTP MCP session in place after the upstream answered HTTP 404 +// (mcp-go: transport.ErrSessionTerminated) for a session id it no longer +// knows, without tearing down the connection or touching connectionEpoch. + +// SessionSnapshot reports the Streamable HTTP transport's current session id +// (empty when the transport has none, e.g. mcp-go cleared it after a 404, or +// when the transport is stdio/SSE) and whether the negotiated protocol is the +// stateless 2026-07-28 era, where sessions do not exist and no re-init applies. +func (c *Client) SessionSnapshot() (id string, modern bool) { + c.mu.RLock() + cl := c.client + info := c.serverInfo + c.mu.RUnlock() + + if info != nil { + modern = mcp.IsModernProtocol(info.ProtocolVersion) + } + if cl == nil { + return "", modern + } + if sh, ok := cl.GetTransport().(*transport.StreamableHTTP); ok { + return sh.GetSessionId(), modern + } + return "", modern +} + +// ReinitializeSession performs a fresh initialize + notifications/initialized +// handshake on the existing mcp-go client and transport, so the upstream +// issues a new session id. It does not Close/Start the transport and does not +// change the connection generation. The legacy-era pin (Spec 058 FR-027) +// applies exactly as on the first handshake. +func (c *Client) ReinitializeSession(ctx context.Context) error { + c.mu.RLock() + cl := c.client + c.mu.RUnlock() + if cl == nil || !c.IsConnected() { + return fmt.Errorf("client not connected") + } + + req := mcp.InitializeRequest{} + req.Params.ProtocolVersion = mcp.LATEST_LEGACY_PROTOCOL_VERSION + req.Params.ClientInfo = mcp.Implementation{Name: "mcpproxy-go", Version: "1.0.0"} + req.Params.Capabilities = mcp.ClientCapabilities{} + + res, err := cl.Initialize(ctx, req) + if err != nil { + return fmt.Errorf("session re-initialize failed: %w", err) + } + if mcp.IsModernProtocol(res.ProtocolVersion) { + return fmt.Errorf("upstream answered MCP protocol version %s on session re-initialize (spec 058 FR-027)", res.ProtocolVersion) + } + + c.mu.Lock() + c.serverInfo = res + c.mu.Unlock() + c.logger.Debug("MCP session re-initialized", zap.String("server", c.config.Name)) + return nil +} diff --git a/internal/upstream/core/session_reinit_mcpgo_test.go b/internal/upstream/core/session_reinit_mcpgo_test.go new file mode 100644 index 000000000..0ff2f60a0 --- /dev/null +++ b/internal/upstream/core/session_reinit_mcpgo_test.go @@ -0,0 +1,171 @@ +package core + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "net/http/httptest" + "sync" + "testing" + "time" + + "github.com/mark3labs/mcp-go/client" + "github.com/mark3labs/mcp-go/client/transport" + "github.com/mark3labs/mcp-go/mcp" + "github.com/stretchr/testify/assert" + "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" +) + +// sessionUpstream is a minimal legacy-era Streamable HTTP upstream with real +// session semantics: initialize issues a new session id; every other request +// must carry a live one or it is answered 404 before anything executes. +type sessionUpstream struct { + mu sync.Mutex + seq int + live map[string]bool + inits int +} + +func newSessionUpstream() *sessionUpstream { return &sessionUpstream{live: map[string]bool{}} } + +// forget makes the upstream drop every session it has issued (server restart). +func (u *sessionUpstream) forget() { + u.mu.Lock() + defer u.mu.Unlock() + u.live = map[string]bool{} +} + +func (u *sessionUpstream) handler() http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + w.WriteHeader(http.StatusMethodNotAllowed) + return + } + body, _ := io.ReadAll(r.Body) + var req struct { + ID json.RawMessage `json:"id"` + Method string `json:"method"` + Params struct { + ProtocolVersion string `json:"protocolVersion"` + } `json:"params"` + } + _ = json.Unmarshal(body, &req) + write := func(res map[string]any) { + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(map[string]any{"jsonrpc": "2.0", "id": json.RawMessage(req.ID), "result": res}) + } + if req.Method == "initialize" { + u.mu.Lock() + u.inits++ + u.seq++ + sid := fmt.Sprintf("sess-%08d", u.seq) + u.live[sid] = true + u.mu.Unlock() + w.Header().Set("Mcp-Session-Id", sid) + write(map[string]any{"protocolVersion": req.Params.ProtocolVersion, "capabilities": map[string]any{"tools": map[string]any{}}, "serverInfo": map[string]any{"name": "sess", "version": "1"}}) + return + } + u.mu.Lock() + ok := u.live[r.Header.Get("Mcp-Session-Id")] + u.mu.Unlock() + if !ok { + w.WriteHeader(http.StatusNotFound) + return + } + switch req.Method { + case "tools/list": + write(map[string]any{"tools": []any{}}) + default: + if len(req.ID) == 0 { + w.WriteHeader(http.StatusAccepted) + return + } + write(map[string]any{}) + } + }) +} + +// T220 / FR-088: pins the mcp-go v1.0.0 behaviour the re-init design relies on. +func TestMCPGo_SessionTerminated_ReinitializeOnSameTransport(t *testing.T) { + up := newSessionUpstream() + srv := httptest.NewServer(up.handler()) + defer srv.Close() + + tr, err := transport.NewStreamableHTTP(srv.URL) + require.NoError(t, err) + cl := client.NewClient(tr) + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + require.NoError(t, cl.Start(ctx)) + t.Cleanup(func() { _ = cl.Close() }) + + initReq := mcp.InitializeRequest{} + initReq.Params.ProtocolVersion = mcp.LATEST_LEGACY_PROTOCOL_VERSION + initReq.Params.ClientInfo = mcp.Implementation{Name: "t", Version: "1"} + _, err = cl.Initialize(ctx, initReq) + require.NoError(t, err) + first := tr.GetSessionId() + require.NotEmpty(t, first) + + // 404 on a non-initialize POST -> ErrSessionTerminated, session id cleared. + up.forget() + _, err = cl.ListTools(ctx, mcp.ListToolsRequest{}) + require.Error(t, err) + assert.True(t, errors.Is(err, transport.ErrSessionTerminated), "got %v", err) + assert.Empty(t, tr.GetSessionId(), "mcp-go clears the session id on 404") + + // 404 with no session id at all is also ErrSessionTerminated (the next + // request goes out without a header, which this upstream rejects). + _, err = cl.ListTools(ctx, mcp.ListToolsRequest{}) + require.Error(t, err) + assert.True(t, errors.Is(err, transport.ErrSessionTerminated), "got %v", err) + + // A second Initialize on the same client/transport stores the new id. + _, err = cl.Initialize(ctx, initReq) + require.NoError(t, err) + second := tr.GetSessionId() + require.NotEmpty(t, second) + assert.NotEqual(t, first, second) + _, err = cl.ListTools(ctx, mcp.ListToolsRequest{}) + require.NoError(t, err) +} + +func TestCoreSessionSnapshotAndReinitialize(t *testing.T) { + disableOAuthForTest(t) + up := newSessionUpstream() + srv := httptest.NewServer(up.handler()) + defer srv.Close() + + cfg := &config.ServerConfig{Name: "sess", Protocol: "streamable-http", URL: srv.URL, Enabled: true} + c, err := NewClient("sess", cfg, zap.NewNop(), nil, nil, nil, secret.NewResolver()) + require.NoError(t, err) + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + require.NoError(t, c.Connect(ctx)) + t.Cleanup(func() { _ = c.Disconnect() }) + + id, modern := c.SessionSnapshot() + assert.NotEmpty(t, id) + assert.False(t, modern) + + up.forget() + _, err = c.ListTools(ctx) + require.Error(t, err) + assert.True(t, errors.Is(err, transport.ErrSessionTerminated), "got %v", err) + id2, _ := c.SessionSnapshot() + assert.Empty(t, id2) + + require.NoError(t, c.ReinitializeSession(ctx)) + id3, _ := c.SessionSnapshot() + assert.NotEmpty(t, id3) + assert.NotEqual(t, id, id3) + _, err = c.ListTools(ctx) + require.NoError(t, err) +} diff --git a/internal/upstream/managed/client.go b/internal/upstream/managed/client.go index f208bcafd..5c3ff9483 100644 --- a/internal/upstream/managed/client.go +++ b/internal/upstream/managed/client.go @@ -105,6 +105,9 @@ type Client struct { // (hand-constructed clients in tests). toolInvoker toolCaller + // sess is the Spec 113-e session re-init state (single-flight, baseline, counter). + sess sessionReinit + // ambiguousProbeInFlight gates the async liveness probe fired after an // ambiguous tools/call cancellation (GH #965) so a burst of canceled calls // results in at most one probe against the upstream. @@ -579,6 +582,10 @@ func (mc *Client) Connect(ctx context.Context) error { mc.connectionEpoch.Store(nextConnectionEpoch()) mc.epochMu.Unlock() + // Spec 113-e FR-080: remember the session id this connect obtained, so a + // transport that loses it before the first managed request can recover. + _, _, _ = mc.sessionState() + // Transition to ready state only if not already ready if mc.StateManager.GetState() != types.StateReady { mc.StateManager.TransitionTo(types.StateReady) @@ -838,6 +845,8 @@ func (mc *Client) GetConnectionStatus() map[string]interface{} { "should_retry": mc.ShouldRetry(), "retry_count": info.RetryCount, "server_name": info.ServerName, + // Spec 113-e FR-085: in-place Streamable HTTP session re-initializations. + "session_reinit_count": mc.SessionReinitCount(), } if info.LastError != nil { @@ -1062,7 +1071,7 @@ func (mc *Client) runListToolsAsLeader(listCtx context.Context, release func() b }() listEpoch := mc.connectionEpoch.Load() - tools, err := mc.coreClient.ListTools(listCtx) + tools, err := mc.listToolsUpstream(listCtx) mc.publishListToolsResult(tools, err) if err != nil { @@ -1201,8 +1210,17 @@ func (mc *Client) callTool(ctx context.Context, toolName string, args map[string // 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(core.WithConnectionGeneration(ctx, callEpoch), toolName, args) + result, err := mc.callToolWithSession(core.WithConnectionGeneration(ctx, callEpoch), invoker, toolName, args, expectedEpoch) if err != nil { + if errors.Is(err, ErrSessionReestablished) || errors.Is(err, ErrConnectionGenerationChanged) { + // Spec 113-e: the session was re-established and the call was + // deliberately not repeated. Not a connection failure. + mc.logger.Warn("Tool call not repeated after upstream session re-initialize", + zap.String("server", mc.GetConfig().Name), + zap.String("tool", toolName), + zap.Error(err)) + return nil, err + } mc.recordCallToolOAuthSignal(toolName, err) // A 429 answered to a tools/call is the same instruction as one answered // to connect (#1040). Recording it here does NOT mark the server @@ -1752,7 +1770,7 @@ func (mc *Client) performHealthCheck() { // 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) + err := mc.probeLiveness(ctx, prober) if err != nil { // Pick up any rate-limit hint this ping's response carried, BEFORE the @@ -2480,7 +2498,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) + tools, err := mc.listToolsUpstream(listCtx) mc.publishListToolsResult(tools, err) if err != nil { mc.logger.Debug("Tool count fetch failed, returning cached value", diff --git a/internal/upstream/managed/prompts.go b/internal/upstream/managed/prompts.go index 34b51ed18..49ce8160b 100644 --- a/internal/upstream/managed/prompts.go +++ b/internal/upstream/managed/prompts.go @@ -14,7 +14,11 @@ func (mc *Client) ListPrompts(ctx context.Context) ([]mcp.Prompt, error) { return nil, fmt.Errorf("client not connected (state: %s)", mc.StateManager.GetState().String()) } - prompts, err := mc.coreClient.ListPrompts(ctx) + var prompts []mcp.Prompt + err := mc.withSession(ctx, string(mcp.MethodPromptsList), func() (e error) { + prompts, e = mc.coreClient.ListPrompts(ctx) + return e + }) if err != nil { mc.logger.Error("Failed to list prompts", zap.String("server", mc.GetConfig().Name), @@ -33,7 +37,11 @@ func (mc *Client) GetPrompt(ctx context.Context, name string, args map[string]st } promptEpoch := mc.connectionEpoch.Load() - result, err := mc.coreClient.GetPrompt(ctx, name, args) + var result *mcp.GetPromptResult + err := mc.withSession(ctx, string(mcp.MethodPromptsGet), func() (e error) { + result, e = mc.coreClient.GetPrompt(ctx, name, args) + return e + }) if err != nil { if mc.isConnectionError(err) { // Guarded: concurrent failures on one dead transport mark it diff --git a/internal/upstream/managed/session_reinit.go b/internal/upstream/managed/session_reinit.go new file mode 100644 index 000000000..c9ac12c1a --- /dev/null +++ b/internal/upstream/managed/session_reinit.go @@ -0,0 +1,346 @@ +package managed + +import ( + "context" + "errors" + "fmt" + "sync" + "sync/atomic" + "time" + + "github.com/mark3labs/mcp-go/client/transport" + "github.com/mark3labs/mcp-go/mcp" + "go.uber.org/zap" + + "github.com/smart-mcp-proxy/mcpproxy-go/internal/config" +) + +// Spec 113-e (G8): when a Streamable HTTP upstream answers HTTP 404 for a +// session id it no longer knows (mcp-go: transport.ErrSessionTerminated), the +// managed client re-initializes the MCP session in place and retries the +// request once, instead of flipping the whole server into Error and tearing +// the connection down. The Streamable HTTP rule that an unknown session is +// rejected before the method executes is what makes a retry safe; because +// mcp-go maps EVERY 404 to ErrSessionTerminated, tools/call is additionally +// retried only for read-only tools whose identity hash is unchanged (FR-081). + +// ErrSessionReestablished is returned for a tools/call that hit a terminated +// session when the call was NOT repeated (write/destructive tool, unknown or +// changed tool identity). The session itself was re-established, so the next +// call works. It is not a connection failure and never marks the server Error. +var ErrSessionReestablished = errors.New("upstream session was lost and re-established; the call was not repeated") + +// reinitTimeout bounds one re-init flight (initialize + tools/list). The +// flight is detached from any single caller's context so one caller going +// away cannot fail the callers waiting on it. +const reinitTimeout = 30 * time.Second + +type reinitFlight struct { + done chan struct{} + err error +} + +type toolIdentity struct { + hash string + readOnly bool +} + +// sessionReinit is the per-client re-init state. The zero value is ready. +type sessionReinit struct { + mu sync.Mutex + known string // last session id seen on the transport + flight *reinitFlight + baseline map[string]toolIdentity // raw tool name -> identity from the last good listing + count atomic.Int64 +} + +// SessionReinitCount reports how many in-place session re-initializations this +// client has performed (FR-085). +func (mc *Client) SessionReinitCount() int64 { return mc.sess.count.Load() } + +func shortID(id string) string { + if len(id) > 8 { + return id[:8] + } + return id +} + +// recordToolBaseline remembers the identity hash and read-only-ness of every +// tool in a successful listing, the reference a post-re-init listing is +// compared with. +func (mc *Client) recordToolBaseline(tools []*config.ToolMetadata) { + next := make(map[string]toolIdentity, len(tools)) + for _, t := range tools { + if t == nil { + continue + } + next[t.Name] = toolIdentity{hash: t.Hash, readOnly: isReadOnlyTool(t)} + } + mc.sess.mu.Lock() + mc.sess.baseline = next + mc.sess.mu.Unlock() +} + +func isReadOnlyTool(t *config.ToolMetadata) bool { + a := t.Annotations + if a == nil || a.ReadOnlyHint == nil || !*a.ReadOnlyHint { + return false + } + return a.DestructiveHint == nil || !*a.DestructiveHint +} + +func (mc *Client) toolIdentityOf(name string) (toolIdentity, bool) { + mc.sess.mu.Lock() + defer mc.sess.mu.Unlock() + id, ok := mc.sess.baseline[name] + return id, ok +} + +// sessionState returns the session id a request would use now, whether the +// negotiated protocol is stateless, and the last known session id. +func (mc *Client) sessionState() (current string, known string, applicable bool) { + if mc.coreClient == nil { + return "", "", false + } + id, modern := mc.coreClient.SessionSnapshot() + if modern { + return "", "", false + } + mc.sess.mu.Lock() + defer mc.sess.mu.Unlock() + if id != "" { + mc.sess.known = id + } + return id, mc.sess.known, true +} + +// reinitSession re-initializes the session identified by stale (the id the +// failing request used, or the known id when the transport had none). It is +// single-flight: concurrent callers for the same stale id share one flight, +// and a caller that finds a newer session already in place does nothing. +func (mc *Client) reinitSession(ctx context.Context, stale, method string) error { + s := &mc.sess + s.mu.Lock() + // A running flight wins over the known-id comparison: initialize installs + // the new session id before the flight's tools/list has been verified, so + // "newer session in place" must not let a caller slip past the flight + // (FR-082, FR-083a). + if f := s.flight; f != nil { + s.mu.Unlock() + return waitFlight(ctx, f) + } + if s.known != stale { + s.mu.Unlock() + return nil + } + f := &reinitFlight{done: make(chan struct{})} + s.flight = f + s.mu.Unlock() + + fctx, cancel := context.WithTimeout(context.WithoutCancel(ctx), reinitTimeout) + defer cancel() + err := mc.runReinit(fctx) + + newID, _ := mc.coreClient.SessionSnapshot() + s.mu.Lock() + if newID != "" || err == nil { + s.known = newID + } + s.flight = nil + s.mu.Unlock() + f.err = err + close(f.done) + + serverName := mc.GetConfig().Name + if err != nil { + mc.logger.Warn("Upstream session re-initialize failed", + zap.String("server", serverName), zap.String("method", method), + zap.String("stale_session", shortID(stale)), zap.Error(err)) + return err + } + s.count.Add(1) + mc.logger.Info("Upstream session terminated (HTTP 404); re-initialized in place", + zap.String("server", serverName), zap.String("method", method), + zap.String("stale_session", shortID(stale)), zap.String("new_session", shortID(newID)), + zap.Int64("reinit_count", s.count.Load())) + return nil +} + +func waitFlight(ctx context.Context, f *reinitFlight) error { + select { + case <-f.done: + return f.err + case <-ctx.Done(): + return ctx.Err() + } +} + +// runReinit is the flight body: initialize on the same transport, then a +// synchronous tools/list to refresh the identity baseline. A changed toolset +// schedules the normal discovery path (differential update / quarantine +// semantics); connectionEpoch is never touched here (FR-083). +func (mc *Client) runReinit(ctx context.Context) error { + if err := mc.coreClient.ReinitializeSession(ctx); err != nil { + return err + } + tools, err := mc.coreClient.ListTools(ctx) + if err != nil { + return fmt.Errorf("tools/list after session re-initialize: %w", err) + } + mc.sess.mu.Lock() + prev := mc.sess.baseline + mc.sess.mu.Unlock() + changed := len(prev) != len(tools) + for _, t := range tools { + if p, ok := prev[t.Name]; !ok || p.hash != t.Hash { + changed = true + } + } + mc.recordToolBaseline(tools) + if changed { + mc.scheduleToolRefresh() + } + return nil +} + +func (mc *Client) scheduleToolRefresh() { + mc.mu.RLock() + cb := mc.toolDiscoveryCallback + mc.mu.RUnlock() + if cb == nil { + return + } + name := mc.GetConfig().Name + go func() { + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + if err := cb(ctx, name); err != nil { + mc.logger.Warn("Tool refresh after session re-initialize failed", + zap.String("server", name), zap.Error(err)) + } + }() +} + +// joinGap re-initializes before sending when the transport has no session id +// but one is known: another request already hit the 404 and mcp-go cleared it, +// so sending now would go out without a session (FR-080, FR-082). +func (mc *Client) joinGap(ctx context.Context, method string) error { + // A flight in progress (initialize done, tools/list not yet verified) is + // joined even though the transport already has a session id again. + mc.sess.mu.Lock() + f := mc.sess.flight + mc.sess.mu.Unlock() + if f != nil { + return waitFlight(ctx, f) + } + cur, known, ok := mc.sessionState() + if !ok || cur != "" || known == "" { + return nil + } + return mc.reinitSession(ctx, known, method) +} + +// withSession runs an idempotent request (list, get, ping) with one re-init and +// one retry on ErrSessionTerminated. Only the final error is returned. +func (mc *Client) withSession(ctx context.Context, method string, fn func() error) error { + if err := mc.joinGap(ctx, method); err != nil { + return err + } + cur, known, ok := mc.sessionState() + stale := cur + if stale == "" { + stale = known + } + err := fn() + if err == nil || !ok || stale == "" || !errors.Is(err, transport.ErrSessionTerminated) { + return err + } + if rerr := mc.reinitSession(ctx, stale, method); rerr != nil { + // Keep the original 404 (so the existing connection-error path still + // matches) and surface why recovery failed (e.g. an auth error). + return errors.Join(err, rerr) + } + return fn() +} + +// probeLiveness is the health ping with session recovery. +func (mc *Client) probeLiveness(ctx context.Context, p livenessProber) error { + if p != livenessProber(mc.coreClient) { + return p.Ping(ctx) + } + return mc.withSession(ctx, string(mcp.MethodPing), func() error { return p.Ping(ctx) }) +} + +// listToolsUpstream is the upstream tools/list with session recovery. +func (mc *Client) listToolsUpstream(ctx context.Context) (tools []*config.ToolMetadata, err error) { + err = mc.withSession(ctx, string(mcp.MethodToolsList), func() error { + var e error + tools, e = mc.coreClient.ListTools(ctx) + return e + }) + if err == nil { + mc.recordToolBaseline(tools) + } + return tools, err +} + +// callToolWithSession is tools/call with session recovery (FR-081). invoker is +// the dispatch surface callTool already chose; recovery applies only when it is +// the real core client (test fakes bypass it). +func (mc *Client) callToolWithSession(ctx context.Context, invoker toolCaller, toolName string, args map[string]interface{}, expectedEpoch *int64) (*mcp.CallToolResult, error) { + if mc.coreClient == nil || invoker != toolCaller(mc.coreClient) { + return invoker.CallTool(ctx, toolName, args) + } + // The identity this call is held against, read before anything can re-init. + pre, preKnown := mc.toolIdentityOf(toolName) + + if err := mc.joinGap(ctx, string(mcp.MethodToolsCall)); err != nil { + return nil, err + } + if refusal := mc.identityRefusal(toolName, pre, preKnown, expectedEpoch); refusal != nil { + return nil, refusal + } + + cur, known, ok := mc.sessionState() + stale := cur + if stale == "" { + stale = known + } + result, err := invoker.CallTool(ctx, toolName, args) + if err == nil || !ok || stale == "" || !errors.Is(err, transport.ErrSessionTerminated) { + return result, err + } + if rerr := mc.reinitSession(ctx, stale, string(mcp.MethodToolsCall)); rerr != nil { + return nil, errors.Join(err, rerr) + } + post, postKnown := mc.toolIdentityOf(toolName) + if !preKnown || !postKnown || post.hash != pre.hash { + if expectedEpoch != nil { + return nil, ErrConnectionGenerationChanged + } + return nil, fmt.Errorf("%w (tool %q identity not confirmed after re-initialize)", ErrSessionReestablished, toolName) + } + // Both the listing the caller was certified against and the re-listed one + // must say read-only: annotations are not part of the identity hash. + if !pre.readOnly || !post.readOnly { + return nil, fmt.Errorf("%w (tool %q is not read-only)", ErrSessionReestablished, toolName) + } + if expectedEpoch != nil && !mc.generationIs(*expectedEpoch) { + return nil, ErrConnectionGenerationChanged + } + return invoker.CallTool(ctx, toolName, args) +} + +// identityRefusal re-runs, after waiting on a re-init flight, the generation +// checks a fresh call would (FR-083a): the pinned epoch and the tool's +// identity hash against the re-listed toolset. +func (mc *Client) identityRefusal(toolName string, pre toolIdentity, preKnown bool, expectedEpoch *int64) error { + if expectedEpoch != nil && !mc.generationIs(*expectedEpoch) { + return ErrConnectionGenerationChanged + } + post, postKnown := mc.toolIdentityOf(toolName) + if preKnown && (!postKnown || post.hash != pre.hash) { + return ErrConnectionGenerationChanged + } + return nil +} diff --git a/internal/upstream/managed/session_reinit_test.go b/internal/upstream/managed/session_reinit_test.go new file mode 100644 index 000000000..8441ce58d --- /dev/null +++ b/internal/upstream/managed/session_reinit_test.go @@ -0,0 +1,356 @@ +package managed + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "net/http/httptest" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "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" +) + +// sessUpstream is a legacy-era Streamable HTTP upstream with real session +// semantics (Spec 113-e): initialize issues a session id; any other request +// without a live id is answered 404 before executing. +type sessUpstream struct { + mu sync.Mutex + seq int + live map[string]bool + inits int + noSession int // non-initialize requests that arrived with no session id + executed map[string]int + readDesc string // description of read_thing; changing it changes its hash + rejectAfter bool // when set, every non-initialize request is 404 (even on a fresh session) + notReadOnly bool // when set, read_thing is listed without readOnlyHint + listEntered chan struct{} // when non-nil, tools/list signals here then blocks on listRelease + listRelease chan struct{} +} + +func newSessUpstream() *sessUpstream { + return &sessUpstream{live: map[string]bool{}, executed: map[string]int{}, readDesc: "reads"} +} + +func (u *sessUpstream) forget() { + u.mu.Lock() + defer u.mu.Unlock() + u.live = map[string]bool{} +} + +func (u *sessUpstream) snapshot() (inits, noSession int, executed map[string]int) { + u.mu.Lock() + defer u.mu.Unlock() + ex := map[string]int{} + for k, v := range u.executed { + ex[k] = v + } + return u.inits, u.noSession, ex +} + +func (u *sessUpstream) handler() http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + w.WriteHeader(http.StatusMethodNotAllowed) + return + } + body, _ := io.ReadAll(r.Body) + var req struct { + ID json.RawMessage `json:"id"` + Method string `json:"method"` + Params struct { + ProtocolVersion string `json:"protocolVersion"` + Name string `json:"name"` + } `json:"params"` + } + _ = json.Unmarshal(body, &req) + write := func(res map[string]any) { + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(map[string]any{"jsonrpc": "2.0", "id": json.RawMessage(req.ID), "result": res}) + } + if req.Method == "initialize" { + u.mu.Lock() + u.inits++ + u.seq++ + sid := fmt.Sprintf("sess-%08d", u.seq) + u.live[sid] = true + u.mu.Unlock() + w.Header().Set("Mcp-Session-Id", sid) + write(map[string]any{"protocolVersion": req.Params.ProtocolVersion, "capabilities": map[string]any{"tools": map[string]any{}}, "serverInfo": map[string]any{"name": "sess", "version": "1"}}) + return + } + sid := r.Header.Get("Mcp-Session-Id") + u.mu.Lock() + ok := u.live[sid] && !u.rejectAfter + if sid == "" { + u.noSession++ + } + desc := u.readDesc + notRO := u.notReadOnly + entered, release := u.listEntered, u.listRelease + u.mu.Unlock() + if !ok { + w.WriteHeader(http.StatusNotFound) + return + } + switch req.Method { + case "tools/call": + u.mu.Lock() + u.executed[req.Params.Name]++ + u.mu.Unlock() + write(map[string]any{"content": []any{map[string]any{"type": "text", "text": "ok"}}}) + case "tools/list": + if entered != nil { + entered <- struct{}{} + <-release + } + ann := map[string]any{"readOnlyHint": !notRO} + write(map[string]any{"tools": []any{ + map[string]any{"name": "read_thing", "description": desc, "inputSchema": map[string]any{"type": "object"}, "annotations": ann}, + map[string]any{"name": "write_thing", "description": "writes", "inputSchema": map[string]any{"type": "object"}}, + }}) + default: + if len(req.ID) == 0 { + w.WriteHeader(http.StatusAccepted) + return + } + write(map[string]any{}) + } + }) +} + +func newSessionClient(t *testing.T) (*Client, *sessUpstream) { + t.Helper() + t.Setenv("MCPPROXY_DISABLE_OAUTH", "true") + up := newSessUpstream() + srv := httptest.NewServer(up.handler()) + t.Cleanup(srv.Close) + + cfg := &config.ServerConfig{Name: "sess", Protocol: "streamable-http", URL: srv.URL, Enabled: true} + mc, err := NewClient("sess", cfg, zap.NewNop(), nil, &config.Config{}, nil, secret.NewResolver()) + require.NoError(t, err) + ctx, cancel := context.WithTimeout(context.Background(), 20*time.Second) + defer cancel() + require.NoError(t, mc.Connect(ctx)) + t.Cleanup(func() { _ = mc.Disconnect() }) + // Discovery baseline: the identity hashes a certified call is held against. + _, err = mc.ListTools(ctx) + require.NoError(t, err) + return mc, up +} + +func tctx(t *testing.T) context.Context { + t.Helper() + ctx, cancel := context.WithTimeout(context.Background(), 20*time.Second) + t.Cleanup(cancel) + return ctx +} + +// SC-006: a session forgotten once -> one re-init, the read-only call succeeds +// once, the client stays Ready, nothing was sent without a session id. +func TestSessionReinit_SingleReadCall(t *testing.T) { + mc, up := newSessionClient(t) + up.forget() + + res, err := mc.CallTool(tctx(t), "read_thing", nil) + require.NoError(t, err) + require.NotNil(t, res) + + inits, noSession, executed := up.snapshot() + assert.Equal(t, 2, inits, "exactly one re-initialize") + assert.Equal(t, 1, executed["read_thing"]) + assert.Equal(t, int64(1), mc.SessionReinitCount()) + assert.True(t, mc.IsConnected()) + assert.Equal(t, "Ready", mc.StateManager.GetState().String()) + _ = noSession +} + +// FR-082: 10 concurrent callers on one terminated session share one re-init. +func TestSessionReinit_ConcurrentCallsSingleFlight(t *testing.T) { + mc, up := newSessionClient(t) + up.forget() + + const n = 10 + var wg sync.WaitGroup + var failures atomic.Int32 + ctx := tctx(t) + for i := 0; i < n; i++ { + wg.Add(1) + go func() { + defer wg.Done() + if _, err := mc.CallTool(ctx, "read_thing", nil); err != nil { + failures.Add(1) + t.Logf("call error: %v", err) + } + }() + } + wg.Wait() + + inits, _, executed := up.snapshot() + assert.Zero(t, failures.Load()) + assert.Equal(t, 2, inits, "10 concurrent callers must trigger exactly one re-initialize") + assert.Equal(t, n, executed["read_thing"]) + assert.Equal(t, int64(1), mc.SessionReinitCount()) + assert.Equal(t, "Ready", mc.StateManager.GetState().String()) +} + +// FR-081: a write tool is not repeated; the session is re-established and the +// next call works. No Error state. +func TestSessionReinit_WriteToolNotRepeated(t *testing.T) { + mc, up := newSessionClient(t) + up.forget() + + _, err := mc.CallTool(tctx(t), "write_thing", nil) + require.Error(t, err) + assert.True(t, errors.Is(err, ErrSessionReestablished), "got %v", err) + + _, _, executed := up.snapshot() + assert.Zero(t, executed["write_thing"], "the write must not be repeated") + assert.Equal(t, "Ready", mc.StateManager.GetState().String()) + + _, err = mc.CallTool(tctx(t), "write_thing", nil) + require.NoError(t, err) + inits, _, executed := up.snapshot() + assert.Equal(t, 2, inits) + assert.Equal(t, 1, executed["write_thing"]) +} + +// FR-081(b)/FR-083a: a read-only tool whose identity hash changed across the +// re-init is not retried. +func TestSessionReinit_ChangedToolHashNotRetried(t *testing.T) { + mc, up := newSessionClient(t) + up.mu.Lock() + up.readDesc = "reads, but now differently" + up.mu.Unlock() + up.forget() + + _, err := mc.CallTool(tctx(t), "read_thing", nil) + require.Error(t, err) + _, _, executed := up.snapshot() + assert.Zero(t, executed["read_thing"]) + + // A pinned call on the same situation is refused with the generation error. + up.forget() + up.mu.Lock() + up.readDesc = "changed again" + up.mu.Unlock() + epoch := mc.ConnectionEpoch() + _, err = mc.CallToolOnEpoch(tctx(t), "read_thing", nil, epoch) + require.Error(t, err) + assert.True(t, errors.Is(err, ErrConnectionGenerationChanged), "got %v", err) +} + +// FR-080/FR-084: a ListTools 404 re-inits, retries once, no Error state. +func TestSessionReinit_ListToolsRetried(t *testing.T) { + mc, up := newSessionClient(t) + up.forget() + + tools, err := mc.ListTools(tctx(t)) + require.NoError(t, err) + assert.Len(t, tools, 2) + assert.Equal(t, int64(1), mc.SessionReinitCount()) + assert.Equal(t, "Ready", mc.StateManager.GetState().String()) +} + +// FR-080: a health ping that hits the 404 first re-inits; the following call +// joins the already-new session and never sends a request without an id. +func TestSessionReinit_PingFirstThenCall(t *testing.T) { + mc, up := newSessionClient(t) + up.forget() + + require.NoError(t, mc.probeLiveness(tctx(t), mc.coreClient)) + _, noSession0, _ := up.snapshot() + + _, err := mc.CallTool(tctx(t), "read_thing", nil) + require.NoError(t, err) + inits, noSession, executed := up.snapshot() + assert.Equal(t, 2, inits) + assert.Equal(t, 1, executed["read_thing"]) + assert.Equal(t, noSession0, noSession, "the call must not go out without a session id") + assert.Equal(t, int64(1), mc.SessionReinitCount()) +} + +// FR-084/FR-086: a retry that is also 404 fails and goes through the existing +// path (the tool is not executed). +func TestSessionReinit_RetryAlso404(t *testing.T) { + mc, up := newSessionClient(t) + up.mu.Lock() + up.rejectAfter = true + up.mu.Unlock() + + _, err := mc.CallTool(tctx(t), "read_thing", nil) + require.Error(t, err) + _, _, executed := up.snapshot() + assert.Zero(t, executed["read_thing"]) +} + +// FR-083: a pinned call succeeds after a re-init and the epoch is unchanged. +func TestSessionReinit_PinnedCallEpochUnchanged(t *testing.T) { + mc, up := newSessionClient(t) + epoch := mc.ConnectionEpoch() + up.forget() + + _, err := mc.CallToolOnEpoch(tctx(t), "read_thing", nil, epoch) + require.NoError(t, err) + assert.Equal(t, epoch, mc.ConnectionEpoch()) + assert.Equal(t, int64(1), mc.SessionReinitCount()) +} + +// Annotations are not part of the identity hash: a tool that stops being +// read-only across the re-init (same hash) must not be repeated. +func TestSessionReinit_ReadOnlyFlippedNotRetried(t *testing.T) { + mc, up := newSessionClient(t) + up.mu.Lock() + up.live = map[string]bool{} + up.notReadOnly = true + up.mu.Unlock() + + _, err := mc.CallTool(tctx(t), "read_thing", nil) + require.Error(t, err) + assert.ErrorIs(t, err, ErrSessionReestablished) + _, _, executed := up.snapshot() + assert.Zero(t, executed["read_thing"], "no longer read-only: must not be executed again") + assert.Equal(t, "Ready", mc.StateManager.GetState().String()) +} + +// FR-082/083a: a caller arriving while a flight is between initialize (new id +// installed) and its verifying tools/list must wait, not send on the new id. +func TestSessionReinit_CallerDuringFlightWaits(t *testing.T) { + mc, up := newSessionClient(t) + up.forget() + entered := make(chan struct{}, 4) + release := make(chan struct{}) + up.mu.Lock() + up.listEntered, up.listRelease = entered, release + up.mu.Unlock() + + ctx := tctx(t) + errs := make(chan error, 2) + go func() { _, err := mc.CallTool(ctx, "read_thing", nil); errs <- err }() + select { + case <-entered: // flight is inside its tools/list; new session id already installed + case <-time.After(10 * time.Second): + t.Fatal("re-init flight never reached tools/list") + } + go func() { _, err := mc.CallTool(ctx, "read_thing", nil); errs <- err }() + time.Sleep(300 * time.Millisecond) + _, _, executed := up.snapshot() + assert.Zero(t, executed["read_thing"], "no call may execute before the flight verified the toolset") + + close(release) + require.NoError(t, <-errs) + require.NoError(t, <-errs) + inits, _, executed := up.snapshot() + assert.Equal(t, 2, inits) + assert.Equal(t, 2, executed["read_thing"]) +}