http.go

  1// SPDX-FileCopyrightText: Amolith <amolith@secluded.site>
  2//
  3// SPDX-License-Identifier: LicenseRef-MutuaL-1.2
  4
  5package mcp
  6
  7import (
  8	"context"
  9	"crypto/subtle"
 10	"errors"
 11	"fmt"
 12	"net"
 13	"net/http"
 14	"strings"
 15	"time"
 16
 17	sdk "github.com/modelcontextprotocol/go-sdk/mcp"
 18)
 19
 20const (
 21	httpReadHeaderTimeout = 10 * time.Second
 22	httpShutdownTimeout   = 5 * time.Second
 23)
 24
 25// RunHTTP serves MCP over Streamable HTTP.
 26func (s *Server) RunHTTP(ctx context.Context, addr, token string) error {
 27	listener, err := new(net.ListenConfig).Listen(ctx, "tcp", addr)
 28	if err != nil {
 29		return fmt.Errorf("listen MCP HTTP address %q: %w", addr, err)
 30	}
 31
 32	return s.serveHTTP(ctx, listener, token)
 33}
 34
 35func (s *Server) serveHTTP(ctx context.Context, listener net.Listener, token string) error {
 36	server := &http.Server{
 37		Handler:           s.HTTPHandler(token),
 38		ReadHeaderTimeout: httpReadHeaderTimeout,
 39	}
 40	serveErr := make(chan error, 1)
 41	go func() {
 42		serveErr <- server.Serve(listener)
 43	}()
 44
 45	select {
 46	case err := <-serveErr:
 47		if errors.Is(err, http.ErrServerClosed) {
 48			return nil
 49		}
 50
 51		return err
 52	case <-ctx.Done():
 53		shutdownCtx, cancel := context.WithTimeout(context.Background(), httpShutdownTimeout)
 54		defer cancel()
 55
 56		if err := server.Shutdown(shutdownCtx); err != nil {
 57			closeErr := server.Close()
 58			serveError := <-serveErr
 59			if serveError != nil && errors.Is(serveError, http.ErrServerClosed) {
 60				serveError = nil
 61			}
 62
 63			return errors.Join(ctx.Err(), err, closeErr, serveError)
 64		}
 65
 66		if err := <-serveErr; err != nil && !errors.Is(err, http.ErrServerClosed) {
 67			return errors.Join(ctx.Err(), err)
 68		}
 69
 70		return ctx.Err()
 71	}
 72}
 73
 74// HTTPHandler returns the authenticated Streamable HTTP MCP handler.
 75func (s *Server) HTTPHandler(token string) http.Handler {
 76	streamable := sdk.NewStreamableHTTPHandler(func(*http.Request) *sdk.Server {
 77		return s.sdk
 78	}, nil)
 79
 80	return newHTTPHandler(token, streamable)
 81}
 82
 83func newHTTPHandler(token string, mcpHandler http.Handler) http.Handler {
 84	mux := http.NewServeMux()
 85	mux.Handle("/mcp", requireBearerToken(token, mcpHandler))
 86
 87	return mux
 88}
 89
 90func requireBearerToken(token string, next http.Handler) http.Handler {
 91	if token == "" {
 92		return next
 93	}
 94
 95	return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
 96		scheme, value, ok := strings.Cut(r.Header.Get("Authorization"), " ")
 97		if !ok || !strings.EqualFold(scheme, "Bearer") ||
 98			subtle.ConstantTimeCompare([]byte(value), []byte(token)) != 1 {
 99			w.Header().Set("WWW-Authenticate", "Bearer")
100			http.Error(w, "unauthorized", http.StatusUnauthorized)
101			return
102		}
103
104		next.ServeHTTP(w, r)
105	})
106}