Skip to content

Commit 1da2ef5

Browse files
piyushbagPiyush Bag
authored andcommitted
mcp: add StreamableHTTPHandler.Close for graceful shutdown
Expose public Close() to tear down all sessions and reject new requests with 503, matching the API proposed in #440.
1 parent 91c010f commit 1da2ef5

2 files changed

Lines changed: 76 additions & 16 deletions

File tree

mcp/streamable.go

Lines changed: 28 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -51,6 +51,7 @@ type StreamableHTTPHandler struct {
5151
onTransportDeletion func(sessionID string) // for testing
5252

5353
mu sync.Mutex
54+
closed bool
5455
sessions map[string]*sessionInfo // keyed by session ID
5556
}
5657

@@ -219,27 +220,30 @@ func NewStreamableHTTPHandler(getServer func(*http.Request) *Server, opts *Strea
219220
return h
220221
}
221222

222-
// closeAll closes all ongoing sessions, for tests.
223-
//
224-
// TODO(rfindley): investigate the best API for callers to configure their
225-
// session lifecycle. (?)
223+
// Close closes the handler, closing and removing all connected sessions,
224+
// and preventing new sessions from being added.
226225
//
227-
// Should we allow passing in a session store? That would allow the handler to
228-
// be stateless.
229-
func (h *StreamableHTTPHandler) closeAll() {
230-
// TODO: if we ever expose this outside of tests, we'll need to do better
231-
// than simply collecting sessions while holding the lock: we need to prevent
232-
// new sessions from being added.
233-
//
234-
// Currently, sessions remove themselves from h.sessions when closed, so we
235-
// can't call Close while holding the lock.
226+
// Close is idempotent.
227+
func (h *StreamableHTTPHandler) Close() error {
236228
h.mu.Lock()
229+
if h.closed {
230+
h.mu.Unlock()
231+
return nil
232+
}
233+
h.closed = true
237234
sessionInfos := slices.Collect(maps.Values(h.sessions))
238-
h.sessions = nil
235+
h.sessions = make(map[string]*sessionInfo)
239236
h.mu.Unlock()
240-
for _, s := range sessionInfos {
241-
s.session.Close()
237+
238+
for _, info := range sessionInfos {
239+
info.session.Close()
242240
}
241+
return nil
242+
}
243+
244+
// closeAll closes all ongoing sessions, for tests.
245+
func (h *StreamableHTTPHandler) closeAll() {
246+
_ = h.Close()
243247
}
244248

245249
// disablelocalhostprotection is a compatibility parameter that allows to disable
@@ -302,6 +306,14 @@ func (h *StreamableHTTPHandler) ServeHTTP(w http.ResponseWriter, req *http.Reque
302306
}
303307
}
304308

309+
h.mu.Lock()
310+
closed := h.closed
311+
h.mu.Unlock()
312+
if closed {
313+
http.Error(w, "handler closed", http.StatusServiceUnavailable)
314+
return
315+
}
316+
305317
// [§2.7] of the spec (2025-06-18): validate the MCP-Protocol-Version
306318
// header. If provided, it must be a supported version. If absent, the
307319
// version is unknown (the request may be an initialize for any version).

mcp/streamable_test.go

Lines changed: 48 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -627,6 +627,54 @@ func TestStreamableServerDisconnect(t *testing.T) {
627627
}
628628
}
629629

630+
func TestStreamableHTTPHandlerClose(t *testing.T) {
631+
server := NewServer(testImpl, nil)
632+
handler := NewStreamableHTTPHandler(func(*http.Request) *Server { return server }, nil)
633+
httpServer := httptest.NewServer(mustNotPanic(t, handler))
634+
defer httpServer.Close()
635+
636+
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
637+
defer cancel()
638+
639+
client := NewClient(testImpl, nil)
640+
clientSession, err := client.Connect(ctx, &StreamableClientTransport{Endpoint: httpServer.URL}, nil)
641+
if err != nil {
642+
t.Fatalf("client.Connect() failed: %v", err)
643+
}
644+
t.Cleanup(func() { _ = clientSession.Close() })
645+
646+
handler.mu.Lock()
647+
if len(handler.sessions) != 1 {
648+
t.Fatalf("want 1 session before Close, got %d", len(handler.sessions))
649+
}
650+
handler.mu.Unlock()
651+
652+
if err := handler.Close(); err != nil {
653+
t.Fatalf("Close() failed: %v", err)
654+
}
655+
if err := handler.Close(); err != nil {
656+
t.Fatalf("second Close() failed: %v", err)
657+
}
658+
659+
handler.mu.Lock()
660+
if len(handler.sessions) != 0 {
661+
t.Fatalf("want 0 sessions after Close, got %d", len(handler.sessions))
662+
}
663+
if !handler.closed {
664+
t.Fatal("want handler.closed true after Close")
665+
}
666+
handler.mu.Unlock()
667+
668+
resp, err := http.Get(httpServer.URL)
669+
if err != nil {
670+
t.Fatalf("http.Get after Close failed: %v", err)
671+
}
672+
defer resp.Body.Close()
673+
if resp.StatusCode != http.StatusServiceUnavailable {
674+
t.Fatalf("got status %d after Close, want %d", resp.StatusCode, http.StatusServiceUnavailable)
675+
}
676+
}
677+
630678
func TestServerTransportCleanup(t *testing.T) {
631679
nClient := 3
632680

0 commit comments

Comments
 (0)