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}