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}