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	"net/http"
 10	"net/http/httptest"
 11	"strings"
 12	"testing"
 13)
 14
 15// server.TRANSPORT.5
 16func TestHTTPHandlerACIDServerTransport5RoutesMCPPath(t *testing.T) {
 17	calls := 0
 18	handler := newHTTPHandler("", http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
 19		calls++
 20		w.WriteHeader(http.StatusNoContent)
 21	}))
 22
 23	recorder := httptest.NewRecorder()
 24	request := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/mcp", nil)
 25	handler.ServeHTTP(recorder, request)
 26
 27	if recorder.Code != http.StatusNoContent {
 28		t.Fatalf("status = %d, want %d", recorder.Code, http.StatusNoContent)
 29	}
 30	if calls != 1 {
 31		t.Fatalf("MCP handler calls = %d, want 1", calls)
 32	}
 33}
 34
 35// server.TRANSPORT.5
 36func TestHTTPHandlerACIDServerTransport5DoesNotRouteOtherPaths(t *testing.T) {
 37	tests := []string{"/", "/mcp/extra"}
 38
 39	for _, path := range tests {
 40		t.Run(path, func(t *testing.T) {
 41			calls := 0
 42			handler := newHTTPHandler("", http.HandlerFunc(func(http.ResponseWriter, *http.Request) {
 43				calls++
 44			}))
 45
 46			recorder := httptest.NewRecorder()
 47			request := httptest.NewRequestWithContext(context.Background(), http.MethodPost, path, nil)
 48			handler.ServeHTTP(recorder, request)
 49
 50			if recorder.Code != http.StatusNotFound {
 51				t.Fatalf("status = %d, want %d", recorder.Code, http.StatusNotFound)
 52			}
 53			if calls != 0 {
 54				t.Fatalf("MCP handler calls = %d, want 0", calls)
 55			}
 56		})
 57	}
 58}
 59
 60// server.SECURITY.7
 61func TestHTTPHandlerACIDServerSecurity7GatesMCPPathWithBearerToken(t *testing.T) {
 62	calls := 0
 63	handler := newHTTPHandler("secret-token", http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
 64		calls++
 65		w.WriteHeader(http.StatusNoContent)
 66	}))
 67
 68	unauthorized := httptest.NewRecorder()
 69	unauthorizedRequest := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/mcp", nil)
 70	handler.ServeHTTP(unauthorized, unauthorizedRequest)
 71
 72	if unauthorized.Code != http.StatusUnauthorized {
 73		t.Fatalf("unauthorized status = %d, want %d", unauthorized.Code, http.StatusUnauthorized)
 74	}
 75	if calls != 0 {
 76		t.Fatalf("MCP handler calls after unauthorized request = %d, want 0", calls)
 77	}
 78
 79	authorized := httptest.NewRecorder()
 80	authorizedRequest := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/mcp", nil)
 81	authorizedRequest.Header.Set("Authorization", "Bearer secret-token")
 82	handler.ServeHTTP(authorized, authorizedRequest)
 83
 84	if authorized.Code != http.StatusNoContent {
 85		t.Fatalf("authorized status = %d, want %d", authorized.Code, http.StatusNoContent)
 86	}
 87	if calls != 1 {
 88		t.Fatalf("MCP handler calls after authorized request = %d, want 1", calls)
 89	}
 90}
 91
 92// server.SECURITY.7
 93func TestServerHTTPHandlerACIDServerSecurity7WrapsSDKHandlerWithBearerAuth(t *testing.T) {
 94	handler := NewServer(nil, "test").HTTPHandler("secret-token")
 95
 96	recorder := httptest.NewRecorder()
 97	request := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/mcp", nil)
 98	handler.ServeHTTP(recorder, request)
 99
100	if recorder.Code != http.StatusUnauthorized {
101		t.Fatalf("status = %d, want %d", recorder.Code, http.StatusUnauthorized)
102	}
103}
104
105// server.SECURITY.3 server.SECURITY.6
106func TestHTTPBearerAuthACIDServerSecurity3And6AllowsMatchingBearerToken(t *testing.T) {
107	calls := 0
108	handler := requireBearerToken("secret-token", http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
109		calls++
110		w.WriteHeader(http.StatusNoContent)
111	}))
112
113	recorder := httptest.NewRecorder()
114	request := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/mcp", nil)
115	request.Header.Set("Authorization", "Bearer secret-token")
116	handler.ServeHTTP(recorder, request)
117
118	if recorder.Code != http.StatusNoContent {
119		t.Fatalf("status = %d, want %d", recorder.Code, http.StatusNoContent)
120	}
121	if calls != 1 {
122		t.Fatalf("next calls = %d, want 1", calls)
123	}
124}
125
126// server.SECURITY.7
127func TestHTTPBearerAuthACIDServerSecurity7RejectsInvalidTokens(t *testing.T) {
128	tests := []struct {
129		name   string
130		header string
131	}{
132		{name: "missing", header: ""},
133		{name: "wrong", header: "Bearer wrong-token"},
134		{name: "basic", header: "Basic secret-token"},
135		{name: "malformed", header: "Bearer"},
136	}
137
138	for _, tt := range tests {
139		t.Run(tt.name, func(t *testing.T) {
140			calls := 0
141			handler := requireBearerToken("secret-token", http.HandlerFunc(func(http.ResponseWriter, *http.Request) {
142				calls++
143			}))
144
145			recorder := httptest.NewRecorder()
146			request := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/mcp", nil)
147			if tt.header != "" {
148				request.Header.Set("Authorization", tt.header)
149			}
150			handler.ServeHTTP(recorder, request)
151
152			if recorder.Code != http.StatusUnauthorized {
153				t.Fatalf("status = %d, want %d", recorder.Code, http.StatusUnauthorized)
154			}
155			if got := recorder.Header().Get("WWW-Authenticate"); got != "Bearer" {
156				t.Fatalf("WWW-Authenticate = %q, want Bearer", got)
157			}
158			if calls != 0 {
159				t.Fatalf("next calls = %d, want 0", calls)
160			}
161		})
162	}
163}
164
165func TestHTTPBearerAuthAllowsRequestsWhenNoTokenConfigured(t *testing.T) {
166	calls := 0
167	handler := requireBearerToken("", http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
168		calls++
169		w.WriteHeader(http.StatusNoContent)
170	}))
171
172	recorder := httptest.NewRecorder()
173	request := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/mcp", nil)
174	handler.ServeHTTP(recorder, request)
175
176	if recorder.Code != http.StatusNoContent {
177		t.Fatalf("status = %d, want %d", recorder.Code, http.StatusNoContent)
178	}
179	if calls != 1 {
180		t.Fatalf("next calls = %d, want 1", calls)
181	}
182}
183
184// server.SECURITY.2
185func TestHTTPBearerAuthACIDServerSecurity2DoesNotExposeTokensInUnauthorizedResponse(t *testing.T) {
186	handler := requireBearerToken("secret-token", http.HandlerFunc(func(http.ResponseWriter, *http.Request) {
187		t.Fatal("next handler was called")
188	}))
189
190	recorder := httptest.NewRecorder()
191	request := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/mcp", nil)
192	request.Header.Set("Authorization", "Bearer submitted-token")
193	handler.ServeHTTP(recorder, request)
194
195	body := recorder.Body.String()
196	if strings.Contains(body, "secret-token") {
197		t.Fatalf("response body exposed configured token: %q", body)
198	}
199	if strings.Contains(body, "submitted-token") {
200		t.Fatalf("response body exposed submitted token: %q", body)
201	}
202}