http_test.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	"errors"
 10	"io"
 11	"net"
 12	"net/http"
 13	"net/http/httptest"
 14	"strings"
 15	"testing"
 16	"time"
 17)
 18
 19// server.TRANSPORT.3
 20func TestRunHTTPACIDServerTransport3ReturnsListenErrorForUnavailableConfiguredAddress(t *testing.T) {
 21	listener, err := new(net.ListenConfig).Listen(context.Background(), "tcp", "127.0.0.1:0")
 22	if err != nil {
 23		t.Fatalf("listen occupied address: %v", err)
 24	}
 25	t.Cleanup(func() {
 26		if closeErr := listener.Close(); closeErr != nil {
 27			t.Fatalf("close listener: %v", closeErr)
 28		}
 29	})
 30
 31	err = NewServer(nil, "test").RunHTTP(context.Background(), listener.Addr().String(), "")
 32	if err == nil {
 33		t.Fatal("RunHTTP() error = nil, want listen error")
 34	}
 35	if !strings.Contains(err.Error(), "listen MCP HTTP address") {
 36		t.Fatalf("RunHTTP() error = %v, want listen address error", err)
 37	}
 38}
 39
 40// server.TRANSPORT.3
 41func TestServeHTTPACIDServerTransport3ServesConfiguredListenerUntilContextCanceled(t *testing.T) {
 42	listener, err := new(net.ListenConfig).Listen(context.Background(), "tcp", "127.0.0.1:0")
 43	if err != nil {
 44		t.Fatalf("listen configured address: %v", err)
 45	}
 46
 47	ctx, cancel := context.WithCancel(context.Background())
 48	serveErr := make(chan error, 1)
 49	go func() {
 50		serveErr <- NewServer(nil, "test").serveHTTP(ctx, listener, "")
 51	}()
 52
 53	requestCtx, stopRequest := context.WithTimeout(context.Background(), time.Second)
 54	defer stopRequest()
 55	request, err := http.NewRequestWithContext(
 56		requestCtx,
 57		http.MethodGet,
 58		"http://"+listener.Addr().String()+"/not-mcp",
 59		nil,
 60	)
 61	if err != nil {
 62		t.Fatalf("new request: %v", err)
 63	}
 64
 65	response, err := http.DefaultClient.Do(request)
 66	if err != nil {
 67		cancel()
 68		t.Fatalf("request configured listener: %v", err)
 69	}
 70	if _, err := io.Copy(io.Discard, response.Body); err != nil {
 71		cancel()
 72		t.Fatalf("read response body: %v", err)
 73	}
 74	if err := response.Body.Close(); err != nil {
 75		cancel()
 76		t.Fatalf("close response body: %v", err)
 77	}
 78	if response.StatusCode != http.StatusNotFound {
 79		cancel()
 80		t.Fatalf("status = %d, want %d", response.StatusCode, http.StatusNotFound)
 81	}
 82
 83	cancel()
 84	select {
 85	case err := <-serveErr:
 86		if !errors.Is(err, context.Canceled) {
 87			t.Fatalf("serveHTTP() error = %v, want context.Canceled", err)
 88		}
 89	case <-time.After(time.Second):
 90		t.Fatal("serveHTTP() did not stop after context cancellation")
 91	}
 92}
 93
 94// server.TRANSPORT.5
 95func TestHTTPHandlerACIDServerTransport5RoutesMCPPath(t *testing.T) {
 96	calls := 0
 97	handler := newHTTPHandler("", http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
 98		calls++
 99		w.WriteHeader(http.StatusNoContent)
100	}))
101
102	recorder := httptest.NewRecorder()
103	request := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/mcp", nil)
104	handler.ServeHTTP(recorder, request)
105
106	if recorder.Code != http.StatusNoContent {
107		t.Fatalf("status = %d, want %d", recorder.Code, http.StatusNoContent)
108	}
109	if calls != 1 {
110		t.Fatalf("MCP handler calls = %d, want 1", calls)
111	}
112}
113
114// server.TRANSPORT.5
115func TestHTTPHandlerACIDServerTransport5DoesNotRouteOtherPaths(t *testing.T) {
116	tests := []string{"/", "/mcp/extra"}
117
118	for _, path := range tests {
119		t.Run(path, func(t *testing.T) {
120			calls := 0
121			handler := newHTTPHandler("", http.HandlerFunc(func(http.ResponseWriter, *http.Request) {
122				calls++
123			}))
124
125			recorder := httptest.NewRecorder()
126			request := httptest.NewRequestWithContext(context.Background(), http.MethodPost, path, nil)
127			handler.ServeHTTP(recorder, request)
128
129			if recorder.Code != http.StatusNotFound {
130				t.Fatalf("status = %d, want %d", recorder.Code, http.StatusNotFound)
131			}
132			if calls != 0 {
133				t.Fatalf("MCP handler calls = %d, want 0", calls)
134			}
135		})
136	}
137}
138
139// server.SECURITY.7
140func TestHTTPHandlerACIDServerSecurity7GatesMCPPathWithBearerToken(t *testing.T) {
141	calls := 0
142	handler := newHTTPHandler("secret-token", http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
143		calls++
144		w.WriteHeader(http.StatusNoContent)
145	}))
146
147	unauthorized := httptest.NewRecorder()
148	unauthorizedRequest := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/mcp", nil)
149	handler.ServeHTTP(unauthorized, unauthorizedRequest)
150
151	if unauthorized.Code != http.StatusUnauthorized {
152		t.Fatalf("unauthorized status = %d, want %d", unauthorized.Code, http.StatusUnauthorized)
153	}
154	if calls != 0 {
155		t.Fatalf("MCP handler calls after unauthorized request = %d, want 0", calls)
156	}
157
158	authorized := httptest.NewRecorder()
159	authorizedRequest := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/mcp", nil)
160	authorizedRequest.Header.Set("Authorization", "Bearer secret-token")
161	handler.ServeHTTP(authorized, authorizedRequest)
162
163	if authorized.Code != http.StatusNoContent {
164		t.Fatalf("authorized status = %d, want %d", authorized.Code, http.StatusNoContent)
165	}
166	if calls != 1 {
167		t.Fatalf("MCP handler calls after authorized request = %d, want 1", calls)
168	}
169}
170
171// server.SECURITY.7
172func TestServerHTTPHandlerACIDServerSecurity7WrapsSDKHandlerWithBearerAuth(t *testing.T) {
173	handler := NewServer(nil, "test").HTTPHandler("secret-token")
174
175	recorder := httptest.NewRecorder()
176	request := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/mcp", nil)
177	handler.ServeHTTP(recorder, request)
178
179	if recorder.Code != http.StatusUnauthorized {
180		t.Fatalf("status = %d, want %d", recorder.Code, http.StatusUnauthorized)
181	}
182}
183
184// server.SECURITY.3 server.SECURITY.6
185func TestHTTPBearerAuthACIDServerSecurity3And6AllowsMatchingBearerToken(t *testing.T) {
186	calls := 0
187	handler := requireBearerToken("secret-token", http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
188		calls++
189		w.WriteHeader(http.StatusNoContent)
190	}))
191
192	recorder := httptest.NewRecorder()
193	request := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/mcp", nil)
194	request.Header.Set("Authorization", "Bearer secret-token")
195	handler.ServeHTTP(recorder, request)
196
197	if recorder.Code != http.StatusNoContent {
198		t.Fatalf("status = %d, want %d", recorder.Code, http.StatusNoContent)
199	}
200	if calls != 1 {
201		t.Fatalf("next calls = %d, want 1", calls)
202	}
203}
204
205// server.SECURITY.7
206func TestHTTPBearerAuthACIDServerSecurity7RejectsInvalidTokens(t *testing.T) {
207	tests := []struct {
208		name   string
209		header string
210	}{
211		{name: "missing", header: ""},
212		{name: "wrong", header: "Bearer wrong-token"},
213		{name: "basic", header: "Basic secret-token"},
214		{name: "malformed", header: "Bearer"},
215	}
216
217	for _, tt := range tests {
218		t.Run(tt.name, func(t *testing.T) {
219			calls := 0
220			handler := requireBearerToken("secret-token", http.HandlerFunc(func(http.ResponseWriter, *http.Request) {
221				calls++
222			}))
223
224			recorder := httptest.NewRecorder()
225			request := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/mcp", nil)
226			if tt.header != "" {
227				request.Header.Set("Authorization", tt.header)
228			}
229			handler.ServeHTTP(recorder, request)
230
231			if recorder.Code != http.StatusUnauthorized {
232				t.Fatalf("status = %d, want %d", recorder.Code, http.StatusUnauthorized)
233			}
234			if got := recorder.Header().Get("WWW-Authenticate"); got != "Bearer" {
235				t.Fatalf("WWW-Authenticate = %q, want Bearer", got)
236			}
237			if calls != 0 {
238				t.Fatalf("next calls = %d, want 0", calls)
239			}
240		})
241	}
242}
243
244func TestHTTPBearerAuthAllowsRequestsWhenNoTokenConfigured(t *testing.T) {
245	calls := 0
246	handler := requireBearerToken("", http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
247		calls++
248		w.WriteHeader(http.StatusNoContent)
249	}))
250
251	recorder := httptest.NewRecorder()
252	request := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/mcp", nil)
253	handler.ServeHTTP(recorder, request)
254
255	if recorder.Code != http.StatusNoContent {
256		t.Fatalf("status = %d, want %d", recorder.Code, http.StatusNoContent)
257	}
258	if calls != 1 {
259		t.Fatalf("next calls = %d, want 1", calls)
260	}
261}
262
263// server.SECURITY.2
264func TestHTTPBearerAuthACIDServerSecurity2DoesNotExposeTokensInUnauthorizedResponse(t *testing.T) {
265	handler := requireBearerToken("secret-token", http.HandlerFunc(func(http.ResponseWriter, *http.Request) {
266		t.Fatal("next handler was called")
267	}))
268
269	recorder := httptest.NewRecorder()
270	request := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/mcp", nil)
271	request.Header.Set("Authorization", "Bearer submitted-token")
272	handler.ServeHTTP(recorder, request)
273
274	body := recorder.Body.String()
275	if strings.Contains(body, "secret-token") {
276		t.Fatalf("response body exposed configured token: %q", body)
277	}
278	if strings.Contains(body, "submitted-token") {
279		t.Fatalf("response body exposed submitted token: %q", body)
280	}
281}