mcp: serve HTTP transport

Amolith created

Change summary

internal/mcp/http.go      | 59 ++++++++++++++++++++++++++++++
internal/mcp/http_test.go | 79 +++++++++++++++++++++++++++++++++++++++++
2 files changed, 138 insertions(+)

Detailed changes

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 {

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