From 1d06a0319a655ecaa4462cbcd975f8b300856d68 Mon Sep 17 00:00:00 2001 From: Amolith Date: Wed, 10 Jun 2026 18:20:19 -0600 Subject: [PATCH] mcp: serve HTTP transport --- internal/mcp/http.go | 59 +++++++++++++++++++++++++++++ internal/mcp/http_test.go | 79 +++++++++++++++++++++++++++++++++++++++ 2 files changed, 138 insertions(+) diff --git a/internal/mcp/http.go b/internal/mcp/http.go index f3864818cc1968313a80c85c0f5b69e0a5e8ef7d..46f7389d3292f42721d5ce6d145888c7da308acf 100644 --- a/internal/mcp/http.go +++ b/internal/mcp/http.go @@ -5,13 +5,72 @@ package mcp import ( + "context" "crypto/subtle" + "errors" + "fmt" + "net" "net/http" "strings" + "time" sdk "github.com/modelcontextprotocol/go-sdk/mcp" ) +const ( + httpReadHeaderTimeout = 10 * time.Second + httpShutdownTimeout = 5 * time.Second +) + +// RunHTTP serves MCP over Streamable HTTP. +func (s *Server) RunHTTP(ctx context.Context, addr, token string) error { + listener, err := new(net.ListenConfig).Listen(ctx, "tcp", addr) + if err != nil { + return fmt.Errorf("listen MCP HTTP address %q: %w", addr, err) + } + + return s.serveHTTP(ctx, listener, token) +} + +func (s *Server) serveHTTP(ctx context.Context, listener net.Listener, token string) error { + server := &http.Server{ + Handler: s.HTTPHandler(token), + ReadHeaderTimeout: httpReadHeaderTimeout, + } + serveErr := make(chan error, 1) + go func() { + serveErr <- server.Serve(listener) + }() + + select { + case err := <-serveErr: + if errors.Is(err, http.ErrServerClosed) { + return nil + } + + return err + case <-ctx.Done(): + shutdownCtx, cancel := context.WithTimeout(context.Background(), httpShutdownTimeout) + defer cancel() + + if err := server.Shutdown(shutdownCtx); err != nil { + closeErr := server.Close() + serveError := <-serveErr + if serveError != nil && errors.Is(serveError, http.ErrServerClosed) { + serveError = nil + } + + return errors.Join(ctx.Err(), err, closeErr, serveError) + } + + if err := <-serveErr; err != nil && !errors.Is(err, http.ErrServerClosed) { + return errors.Join(ctx.Err(), err) + } + + return ctx.Err() + } +} + // HTTPHandler returns the authenticated Streamable HTTP MCP handler. func (s *Server) HTTPHandler(token string) http.Handler { streamable := sdk.NewStreamableHTTPHandler(func(*http.Request) *sdk.Server { diff --git a/internal/mcp/http_test.go b/internal/mcp/http_test.go index d323f31cd90104cc1c2869737e50a9497f82932d..5ffadc86a3ef8bc12dae7393ab76b0ceeb210eb2 100644 --- a/internal/mcp/http_test.go +++ b/internal/mcp/http_test.go @@ -6,12 +6,91 @@ package mcp import ( "context" + "errors" + "io" + "net" "net/http" "net/http/httptest" "strings" "testing" + "time" ) +// server.TRANSPORT.3 +func TestRunHTTPACIDServerTransport3ReturnsListenErrorForUnavailableConfiguredAddress(t *testing.T) { + listener, err := new(net.ListenConfig).Listen(context.Background(), "tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("listen occupied address: %v", err) + } + t.Cleanup(func() { + if closeErr := listener.Close(); closeErr != nil { + t.Fatalf("close listener: %v", closeErr) + } + }) + + err = NewServer(nil, "test").RunHTTP(context.Background(), listener.Addr().String(), "") + if err == nil { + t.Fatal("RunHTTP() error = nil, want listen error") + } + if !strings.Contains(err.Error(), "listen MCP HTTP address") { + t.Fatalf("RunHTTP() error = %v, want listen address error", err) + } +} + +// server.TRANSPORT.3 +func TestServeHTTPACIDServerTransport3ServesConfiguredListenerUntilContextCanceled(t *testing.T) { + listener, err := new(net.ListenConfig).Listen(context.Background(), "tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("listen configured address: %v", err) + } + + ctx, cancel := context.WithCancel(context.Background()) + serveErr := make(chan error, 1) + go func() { + serveErr <- NewServer(nil, "test").serveHTTP(ctx, listener, "") + }() + + requestCtx, stopRequest := context.WithTimeout(context.Background(), time.Second) + defer stopRequest() + request, err := http.NewRequestWithContext( + requestCtx, + http.MethodGet, + "http://"+listener.Addr().String()+"/not-mcp", + nil, + ) + if err != nil { + t.Fatalf("new request: %v", err) + } + + response, err := http.DefaultClient.Do(request) + if err != nil { + cancel() + t.Fatalf("request configured listener: %v", err) + } + if _, err := io.Copy(io.Discard, response.Body); err != nil { + cancel() + t.Fatalf("read response body: %v", err) + } + if err := response.Body.Close(); err != nil { + cancel() + t.Fatalf("close response body: %v", err) + } + if response.StatusCode != http.StatusNotFound { + cancel() + t.Fatalf("status = %d, want %d", response.StatusCode, http.StatusNotFound) + } + + cancel() + select { + case err := <-serveErr: + if !errors.Is(err, context.Canceled) { + t.Fatalf("serveHTTP() error = %v, want context.Canceled", err) + } + case <-time.After(time.Second): + t.Fatal("serveHTTP() did not stop after context cancellation") + } +} + // server.TRANSPORT.5 func TestHTTPHandlerACIDServerTransport5RoutesMCPPath(t *testing.T) { calls := 0