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}