@@ -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 {
@@ -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