This is an automated email from the ASF dual-hosted git repository.
lizhimins pushed a commit to branch rocketmq-studio
in repository https://gitbox.apache.org/repos/asf/rocketmq-dashboard.git
The following commit(s) were added to refs/heads/rocketmq-studio by this push:
new 87c25b56c fix(rmqctl): recover incomplete MCP reinitialization (#5491)
87c25b56c is described below
commit 87c25b56c256dac754abdcb88b90ffd504219967
Author: zmuxuny <[email protected]>
AuthorDate: Fri Oct 9 23:39:42 2026 -0700
fix(rmqctl): recover incomplete MCP reinitialization (#5491)
`mcp_message.go` advanced the generation before `notifications/initialized`
had been sent, so a
concurrent caller could see the session as already rebuilt and send a
request into a half-initialized
session - contradicting the contract `mcp_http.go` states next to it. The
generation now advances only
once initialization completes, and an interrupted initialization is
recoverable.
`go test ./...` green across all seven rmqctl packages.
---
rmqctl/internal/studio/mcp_http.go | 7 +-
rmqctl/internal/studio/mcp_message.go | 45 ++-
rmqctl/internal/studio/reconnect_failure_test.go | 381 +++++++++++++++++++++++
3 files changed, 418 insertions(+), 15 deletions(-)
diff --git a/rmqctl/internal/studio/mcp_http.go
b/rmqctl/internal/studio/mcp_http.go
index 83259e1a0..f8a1bb3af 100644
--- a/rmqctl/internal/studio/mcp_http.go
+++ b/rmqctl/internal/studio/mcp_http.go
@@ -59,9 +59,10 @@ type MCPClientSession struct {
sendMu sync.RWMutex
state struct {
sync.RWMutex
- initialize *mcptransport.JSONRPCRequest
- generation uint64
- ready bool
+ initialize *mcptransport.JSONRPCRequest
+ generation uint64
+ ready bool
+ reconnectPending bool
}
}
diff --git a/rmqctl/internal/studio/mcp_message.go
b/rmqctl/internal/studio/mcp_message.go
index 9acefefc2..8babade41 100644
--- a/rmqctl/internal/studio/mcp_message.go
+++ b/rmqctl/internal/studio/mcp_message.go
@@ -151,18 +151,35 @@ func (session *MCPClientSession) sendNotification(
return err
}
+var errReconnectIncomplete = errors.New("MCP session reinitialization is
incomplete")
+
func (session *MCPClientSession) sendWithReconnect(ctx context.Context,
skipReconnect bool, send func() error) error {
- session.sendMu.RLock()
- generation, hadSession := session.sendSnapshot()
- err := send()
- session.sendMu.RUnlock()
- if err != nil && hadSession && !skipReconnect && errors.Is(err,
mcptransport.ErrSessionTerminated) {
+ var err error
+ // A call may recover once and send at most twice. Only a definitive
session
+ // termination permits replay; an incomplete handshake has not sent the
call.
+ for attempt := 0; attempt < 2; attempt++ {
+ session.sendMu.RLock()
+ generation, hadSession := session.sendSnapshot()
+ session.state.RLock()
+ pending := session.state.reconnectPending
+ session.state.RUnlock()
+ if pending && !skipReconnect {
+ err = errReconnectIncomplete
+ } else {
+ err = send()
+ }
+ session.sendMu.RUnlock()
+ if err == nil || skipReconnect || attempt == 1 {
+ return err
+ }
+ if !pending && !(hadSession && errors.Is(err,
mcptransport.ErrSessionTerminated)) {
+ return err
+ }
if reconnectErr := session.reinitialize(ctx, generation);
reconnectErr != nil {
return reconnectErr
}
- session.sendMu.RLock()
- err = send()
- session.sendMu.RUnlock()
+ // Recheck pending under the send lock on the next attempt:
another
+ // caller may have started a newer reconnect before we regain
the lock.
}
return err
}
@@ -186,6 +203,7 @@ func (session *MCPClientSession) recordInitialization(
}
session.state.Lock()
session.state.initialize = new(request)
+ session.state.generation++
session.state.ready = false
session.state.Unlock()
return nil
@@ -201,9 +219,6 @@ func (session *MCPClientSession)
applyInitialization(response *mcptransport.JSON
if protocolVersion != "" {
session.transport.SetProtocolVersion(protocolVersion)
}
- session.state.Lock()
- session.state.generation++
- session.state.Unlock()
return nil
}
@@ -223,7 +238,7 @@ func (session *MCPClientSession) reinitialize(ctx
context.Context, expectedGener
session.sendMu.Lock()
defer session.sendMu.Unlock()
session.state.RLock()
- if session.state.generation != expectedGeneration {
+ if session.state.generation != expectedGeneration &&
!session.state.reconnectPending {
session.state.RUnlock()
return nil
}
@@ -232,6 +247,10 @@ func (session *MCPClientSession) reinitialize(ctx
context.Context, expectedGener
if initialize == nil {
return fmt.Errorf("MCP session terminated before initialization
could be replayed")
}
+ session.state.Lock()
+ session.state.reconnectPending = true
+ session.state.ready = false
+ session.state.Unlock()
response, err := session.transport.SendRequest(ctx, *initialize)
if err != nil {
return fmt.Errorf("reinitialize MCP session: %w", err)
@@ -254,6 +273,8 @@ func (session *MCPClientSession) reinitialize(ctx
context.Context, expectedGener
return fmt.Errorf("reinitialize MCP session: send initialized
notification: %w", err)
}
session.state.Lock()
+ session.state.generation++
+ session.state.reconnectPending = false
session.state.ready = true
session.state.Unlock()
return nil
diff --git a/rmqctl/internal/studio/reconnect_failure_test.go
b/rmqctl/internal/studio/reconnect_failure_test.go
new file mode 100644
index 000000000..90ff3d759
--- /dev/null
+++ b/rmqctl/internal/studio/reconnect_failure_test.go
@@ -0,0 +1,381 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+package studio
+
+import (
+ "bytes"
+ "context"
+ "encoding/json"
+ "fmt"
+ "net/http"
+ "sync"
+ "sync/atomic"
+ "testing"
+ "time"
+
+ mcptransport "github.com/mark3labs/mcp-go/client/transport"
+ "github.com/mark3labs/mcp-go/mcp"
+)
+
+func TestSessionRetriesIncompleteReconnectBeforeNextRequest(t *testing.T) {
+ initializes, initialized, unreadyCalls := 0, 0, 0
+ ready := false
+ client := NewClient(&http.Client{Transport: roundTripFunc(func(request
*http.Request) (*http.Response, error) {
+ if request.Method == http.MethodDelete {
+ return mcpHTTPResponse(http.StatusNoContent, "", ""),
nil
+ }
+ payload := readMCPPayload(t, request)
+ switch payload.Method {
+ case string(mcp.MethodInitialize):
+ initializes++
+ ready = false
+ response := mcpJSONResultResponse(payload.ID,
map[string]any{
+ "protocolVersion": mcp.LATEST_PROTOCOL_VERSION,
+ "capabilities": map[string]any{},
+ "serverInfo": map[string]any{"name":
"studio", "version": "1"},
+ })
+ response.Header.Set(mcptransport.HeaderKeySessionID,
fmt.Sprintf("session-%d", initializes))
+ return response, nil
+ case string(mcp.MethodNotificationInitialized):
+ initialized++
+ if initialized == 2 {
+ return mcpHTTPResponse(http.StatusBadGateway,
"application/json", `{"code":"UNAVAILABLE","message":"Transient upstream
failure"}`), nil
+ }
+ ready = true
+ return mcpHTTPResponse(http.StatusAccepted, "", ""), nil
+ case string(mcp.MethodToolsList):
+ if request.Header.Get(mcptransport.HeaderKeySessionID)
== "session-1" {
+ return mcpHTTPResponse(http.StatusNotFound,
"text/plain", "Session expired"), nil
+ }
+ if !ready {
+ unreadyCalls++
+ return mcpHTTPResponse(http.StatusOK,
"application/json",
fmt.Sprintf(`{"jsonrpc":"2.0","id":%s,"error":{"code":-32600,"message":"Session
not initialized"}}`, payload.ID)), nil
+ }
+ return mcpJSONResultResponse(payload.ID,
map[string]any{"tools": []any{}}), nil
+ default:
+ return nil, fmt.Errorf("unexpected method: %s",
payload.Method)
+ }
+ })})
+ session := newMCPTestSession(t, client, Target{
+ Server: "http://localhost", InstanceID: "instance-dev",
+ Credential: Credential{AccessKey: "test-ak", SecretKey:
"test-sk"}, Timeout: time.Second,
+ })
+ defer session.Close()
+ ctx := context.Background()
+ initializeClientSession(t, ctx, session)
+ request :=
json.RawMessage(`{"jsonrpc":"2.0","id":"list-1","method":"tools/list","params":{}}`)
+ if _, _, err := session.SendMessage(ctx, request); err == nil {
+ t.Fatal("expected the first call to report the failed reconnect
notification")
+ }
+ response, ok, err := session.SendMessage(ctx, request)
+ t.Logf("retry response=%s ok=%t err=%v initialize=%d initialized=%d
unready-tool-calls=%d", response, ok, err, initializes, initialized,
unreadyCalls)
+ if err != nil || !ok || !bytes.Contains(response, []byte(`"tools":[]`))
|| unreadyCalls != 0 {
+ t.Fatalf("retry did not repair the incomplete handshake before
sending tools/list")
+ }
+}
+
+func TestSessionQueuedCallerDoesNotUseIncompleteReconnect(t *testing.T) {
+ var lock sync.Mutex
+ initializes, initialized, unreadyCalls := 0, 0, 0
+ ready := false
+ var expiredRequests atomic.Int32
+ bothExpired := make(chan struct{})
+ client := NewClient(&http.Client{Transport: roundTripFunc(func(request
*http.Request) (*http.Response, error) {
+ if request.Method == http.MethodDelete {
+ return mcpHTTPResponse(http.StatusNoContent, "", ""),
nil
+ }
+ payload := readMCPPayload(t, request)
+ if payload.Method == string(mcp.MethodToolsList) &&
request.Header.Get(mcptransport.HeaderKeySessionID) == "session-1" {
+ if expiredRequests.Add(1) == 2 {
+ close(bothExpired)
+ }
+ select {
+ case <-bothExpired:
+ case <-request.Context().Done():
+ return nil, request.Context().Err()
+ }
+ return mcpHTTPResponse(http.StatusNotFound,
"text/plain", "Session expired"), nil
+ }
+ lock.Lock()
+ defer lock.Unlock()
+ switch payload.Method {
+ case string(mcp.MethodInitialize):
+ initializes++
+ ready = false
+ response := mcpJSONResultResponse(payload.ID,
map[string]any{
+ "protocolVersion": mcp.LATEST_PROTOCOL_VERSION,
+ "capabilities": map[string]any{},
+ "serverInfo": map[string]any{"name":
"studio", "version": "1"},
+ })
+ response.Header.Set(mcptransport.HeaderKeySessionID,
fmt.Sprintf("session-%d", initializes))
+ return response, nil
+ case string(mcp.MethodNotificationInitialized):
+ initialized++
+ if initialized == 2 {
+ return mcpHTTPResponse(http.StatusBadGateway,
"application/json", `{"code":"UNAVAILABLE"}`), nil
+ }
+ ready = true
+ return mcpHTTPResponse(http.StatusAccepted, "", ""), nil
+ case string(mcp.MethodToolsList):
+ if !ready {
+ unreadyCalls++
+ return mcpHTTPResponse(http.StatusOK,
"application/json",
fmt.Sprintf(`{"jsonrpc":"2.0","id":%s,"error":{"code":-32600,"message":"Session
not initialized"}}`, payload.ID)), nil
+ }
+ return mcpJSONResultResponse(payload.ID,
map[string]any{"tools": []any{}}), nil
+ default:
+ return nil, fmt.Errorf("unexpected method: %s",
payload.Method)
+ }
+ })})
+ session := newMCPTestSession(t, client, Target{
+ Server: "http://localhost", InstanceID: "instance-dev",
+ Credential: Credential{AccessKey: "test-ak", SecretKey:
"test-sk"}, Timeout: 3 * time.Second,
+ })
+ defer session.Close()
+ initializeClientSession(t, context.Background(), session)
+ type result struct {
+ payload json.RawMessage
+ err error
+ }
+ results := make(chan result, 2)
+ for id := 0; id < 2; id++ {
+ go func() {
+ payload, _, err :=
session.SendMessage(context.Background(),
json.RawMessage(fmt.Sprintf(`{"jsonrpc":"2.0","id":%d,"method":"tools/list","params":{}}`,
id+2)))
+ results <- result{payload, err}
+ }()
+ }
+ successes, failures := 0, 0
+ for i := 0; i < 2; i++ {
+ result := <-results
+ if result.err != nil {
+ failures++
+ } else if bytes.Contains(result.payload, []byte(`"tools":[]`)) {
+ successes++
+ }
+ t.Logf("caller result=%s err=%v", result.payload, result.err)
+ }
+ lock.Lock()
+ defer lock.Unlock()
+ t.Logf("initialize=%d initialized=%d unready-tool-calls=%d",
initializes, initialized, unreadyCalls)
+ if successes != 1 || failures != 1 || unreadyCalls != 0 {
+ t.Fatalf("queued caller used the incomplete reconnect:
successes=%d failures=%d unready=%d", successes, failures, unreadyCalls)
+ }
+}
+
+// reconnectFailureTransport isolates failure phases without retrying an
ordinary
+// request that may already have been accepted. The HTTP tests above exercise
the
+// same lifecycle through mcp-go's real Streamable HTTP transport.
+type reconnectFailureTransport struct {
+ mu sync.Mutex
+ sessionID string
+ ready bool
+ initializeFailure string
+ failInitializes int
+ failNotifications int
+ initializeCalls int
+ notificationCalls int
+ ordinaryCalls int
+ unreadyCalls int
+ ordinaryError error
+ ordinaryRPCError bool
+ alwaysExpired bool
+}
+
+func (transport *reconnectFailureTransport) Start(context.Context) error
{ return nil }
+func (transport *reconnectFailureTransport) Close() error
{ return nil }
+func (transport *reconnectFailureTransport)
SetNotificationHandler(func(mcp.JSONRPCNotification)) {}
+func (transport *reconnectFailureTransport) SetProtocolVersion(string)
{}
+func (transport *reconnectFailureTransport) GetSessionId() string {
+ transport.mu.Lock()
+ defer transport.mu.Unlock()
+ return transport.sessionID
+}
+func (transport *reconnectFailureTransport) SendRequest(_ context.Context,
request mcptransport.JSONRPCRequest) (*mcptransport.JSONRPCResponse, error) {
+ transport.mu.Lock()
+ defer transport.mu.Unlock()
+ response := &mcptransport.JSONRPCResponse{JSONRPC: mcp.JSONRPC_VERSION,
ID: request.ID, Result: json.RawMessage(`{"tools":[]}`)}
+ if request.Method == string(mcp.MethodInitialize) {
+ transport.initializeCalls++
+ transport.ready = false
+ if transport.failInitializes > 0 {
+ transport.failInitializes--
+ switch transport.initializeFailure {
+ case "transport":
+ return nil, &MCPHTTPStatusError{StatusCode:
http.StatusBadGateway}
+ case "empty":
+ return nil, nil
+ case "rpc":
+ response.Result = nil
+ response.Error = &mcp.JSONRPCErrorDetails{Code:
-32603, Message: "initialize failed"}
+ return response, nil
+ case "malformed":
+ response.Result =
json.RawMessage(`{"protocolVersion":123}`)
+ return response, nil
+ }
+ }
+ transport.sessionID = "replacement"
+ response.Result = json.RawMessage(`{"protocolVersion":"` +
mcp.LATEST_PROTOCOL_VERSION + `"}`)
+ return response, nil
+ }
+ transport.ordinaryCalls++
+ if transport.sessionID == "expired" || transport.alwaysExpired {
+ transport.sessionID = ""
+ return nil, mcptransport.ErrSessionTerminated
+ }
+ if !transport.ready {
+ transport.unreadyCalls++
+ }
+ if transport.ordinaryError != nil {
+ return nil, transport.ordinaryError
+ }
+ if transport.ordinaryRPCError {
+ response.Result = nil
+ response.Error = &mcp.JSONRPCErrorDetails{Code: -32603,
Message: "tool failed"}
+ }
+ return response, nil
+}
+func (transport *reconnectFailureTransport) SendNotification(_
context.Context, _ mcp.JSONRPCNotification) error {
+ transport.mu.Lock()
+ defer transport.mu.Unlock()
+ transport.notificationCalls++
+ if transport.failNotifications > 0 {
+ transport.failNotifications--
+ return &MCPHTTPStatusError{StatusCode: http.StatusBadGateway}
+ }
+ transport.ready = true
+ return nil
+}
+func failureTestSession(transport *reconnectFailureTransport)
*MCPClientSession {
+ session := &MCPClientSession{transport: transport}
+ session.state.initialize = &mcptransport.JSONRPCRequest{JSONRPC:
mcp.JSONRPC_VERSION, ID: mcp.NewRequestId(1), Method:
string(mcp.MethodInitialize)}
+ session.state.generation = 1
+ session.state.ready = true
+ return session
+}
+func sendFailureTestRequest(session *MCPClientSession) error {
+ _, _, err := session.SendMessage(context.Background(),
json.RawMessage(`{"jsonrpc":"2.0","id":2,"method":"tools/list","params":{}}`))
+ return err
+}
+
+func TestSessionKeepsFailedInitializeRecoverable(t *testing.T) {
+ for _, failure := range []string{"transport", "empty", "rpc",
"malformed"} {
+ t.Run(failure, func(t *testing.T) {
+ transport := &reconnectFailureTransport{sessionID:
"expired", ready: true, initializeFailure: failure, failInitializes: 1}
+ session := failureTestSession(transport)
+ if err := sendFailureTestRequest(session); err == nil {
+ t.Fatal("expected failed initialize")
+ }
+ if !session.state.reconnectPending ||
session.state.generation != 1 {
+ t.Fatalf("failed handshake published state:
pending=%t generation=%d", session.state.reconnectPending,
session.state.generation)
+ }
+ if err := sendFailureTestRequest(session); err != nil {
+ t.Fatalf("later recovery failed: %v", err)
+ }
+ if session.state.reconnectPending ||
session.state.generation != 2 || transport.unreadyCalls != 0 {
+ t.Fatalf("recovery state: pending=%t
generation=%d unready=%d", session.state.reconnectPending,
session.state.generation, transport.unreadyCalls)
+ }
+ })
+ }
+}
+
+func TestSessionRepeatedHandshakeFailureHasBoundedRetries(t *testing.T) {
+ transport := &reconnectFailureTransport{sessionID: "expired", ready:
true, failNotifications: 2}
+ session := failureTestSession(transport)
+ for attempt := 1; attempt <= 2; attempt++ {
+ if err := sendFailureTestRequest(session); err == nil {
+ t.Fatal("expected failed initialized notification")
+ }
+ if transport.initializeCalls != attempt ||
transport.ordinaryCalls != 1 || transport.unreadyCalls != 0 {
+ t.Fatalf("failure retry exceeded budget: initialize=%d
ordinary=%d unready=%d", transport.initializeCalls, transport.ordinaryCalls,
transport.unreadyCalls)
+ }
+ }
+ if err := sendFailureTestRequest(session); err != nil {
+ t.Fatalf("later recovery: %v", err)
+ }
+ if transport.initializeCalls != 3 || transport.ordinaryCalls != 2 {
+ t.Fatalf("unexpected recovery counts: initialize=%d
ordinary=%d", transport.initializeCalls, transport.ordinaryCalls)
+ }
+}
+
+func TestSessionConcurrentCallersSharePendingRecovery(t *testing.T) {
+ transport := &reconnectFailureTransport{sessionID: "expired", ready:
true, failNotifications: 1}
+ session := failureTestSession(transport)
+ if err := sendFailureTestRequest(session); err == nil {
+ t.Fatal("expected initial reconnect failure")
+ }
+ const callers = 12
+ var wait sync.WaitGroup
+ errors := make(chan error, callers)
+ for i := 0; i < callers; i++ {
+ wait.Go(func() { errors <- sendFailureTestRequest(session) })
+ }
+ wait.Wait()
+ close(errors)
+ for err := range errors {
+ if err != nil {
+ t.Errorf("concurrent caller: %v", err)
+ }
+ }
+ if transport.initializeCalls != 2 || transport.notificationCalls != 2
|| transport.unreadyCalls != 0 {
+ t.Fatalf("pending recovery was not shared: initialize=%d
initialized=%d unready=%d", transport.initializeCalls,
transport.notificationCalls, transport.unreadyCalls)
+ }
+}
+
+func TestSessionDoesNotReplayAmbiguousOrdinaryFailures(t *testing.T) {
+ for _, failure := range []string{"timeout", "http", "rpc"} {
+ t.Run(failure, func(t *testing.T) {
+ transport := &reconnectFailureTransport{sessionID:
"ready", ready: true}
+ switch failure {
+ case "timeout":
+ transport.ordinaryError =
context.DeadlineExceeded
+ case "http":
+ transport.ordinaryError =
&MCPHTTPStatusError{StatusCode: http.StatusBadGateway}
+ case "rpc":
+ transport.ordinaryRPCError = true
+ }
+ session := failureTestSession(transport)
+ _, _, _ = session.SendMessage(context.Background(),
json.RawMessage(`{"jsonrpc":"2.0","id":2,"method":"tools/call","params":{"name":"example","arguments":{}}}`))
+ if transport.ordinaryCalls != 1 ||
transport.initializeCalls != 0 {
+ t.Fatalf("ambiguous tool failure was replayed:
ordinary=%d initialize=%d", transport.ordinaryCalls, transport.initializeCalls)
+ }
+ })
+ }
+}
+
+func TestSessionDoesNotLoopOnRepeatedSessionTermination(t *testing.T) {
+ transport := &reconnectFailureTransport{sessionID: "expired", ready:
true, alwaysExpired: true}
+ session := failureTestSession(transport)
+ if err := sendFailureTestRequest(session); err == nil {
+ t.Fatal("expected repeated session termination")
+ }
+ if transport.ordinaryCalls != 2 || transport.initializeCalls != 1 {
+ t.Fatalf("recovery exceeded one attempt: ordinary=%d
initialize=%d", transport.ordinaryCalls, transport.initializeCalls)
+ }
+}
+
+func TestSessionClientInitializeDoesNotClearPendingRecovery(t *testing.T) {
+ transport := &reconnectFailureTransport{sessionID: "replacement"}
+ session := failureTestSession(transport)
+ session.state.reconnectPending = true
+ request := *session.state.initialize
+ response := &mcptransport.JSONRPCResponse{Result:
json.RawMessage(`{"protocolVersion":"` + mcp.LATEST_PROTOCOL_VERSION + `"}`)}
+ if err := session.recordInitialization(request, response); err != nil {
+ t.Fatal(err)
+ }
+ if !session.state.reconnectPending {
+ t.Fatal("client initialize cleared pending recovery before
initialized succeeded")
+ }
+}