|
1 | 1 | package ghmcp |
2 | 2 |
|
3 | 3 | import ( |
| 4 | + "bytes" |
4 | 5 | "context" |
5 | 6 | "encoding/json" |
| 7 | + "fmt" |
| 8 | + "net/http" |
| 9 | + "net/http/httptest" |
6 | 10 | "testing" |
7 | 11 |
|
8 | 12 | "github.com/github/github-mcp-server/v2/internal/oauth" |
@@ -237,6 +241,143 @@ func TestOAuthTypedRegistrationLegacyElicitation(t *testing.T) { |
237 | 241 | } |
238 | 242 | } |
239 | 243 |
|
| 244 | +func TestOAuthTypedHTTPHeaderBoundary(t *testing.T) { |
| 245 | + for _, tc := range []struct { |
| 246 | + name string |
| 247 | + arguments string |
| 248 | + owner string |
| 249 | + repo string |
| 250 | + mismatch bool |
| 251 | + valid bool |
| 252 | + }{ |
| 253 | + {"owner-mismatch", `{"owner":"octo","repo":"hello","mode":"valid"}`, "other", "hello", true, true}, |
| 254 | + {"repo-mismatch", `{"owner":"octo","repo":"hello","mode":"valid"}`, "octo", "other", true, true}, |
| 255 | + {"missing-header", `{"owner":"octo","repo":"hello","mode":"valid"}`, "", "hello", true, true}, |
| 256 | + {"matching", `{"owner":"octo","repo":"hello","mode":"valid"}`, "octo", "hello", false, true}, |
| 257 | + {"missing-required", `{}`, "", "", false, false}, |
| 258 | + {"invalid-enum", `{"owner":"octo","repo":"hello","mode":"invalid"}`, "octo", "hello", false, false}, |
| 259 | + {"invalid-owner-type", `{"owner":42,"repo":"hello","mode":"valid"}`, "42", "hello", false, false}, |
| 260 | + {"null-owner", `{"owner":null,"repo":"hello","mode":"valid"}`, "", "hello", false, false}, |
| 261 | + {"non-object", `[]`, "", "", false, false}, |
| 262 | + } { |
| 263 | + t.Run(tc.name, func(t *testing.T) { |
| 264 | + fake := oauthPendingAuthenticator() |
| 265 | + handlerCalls, receivingCalls, preflightCalls, normalizerCalls := 0, 0, 0, 0 |
| 266 | + type input struct { |
| 267 | + Owner string `json:"owner"` |
| 268 | + Repo string `json:"repo"` |
| 269 | + Mode string `json:"mode"` |
| 270 | + } |
| 271 | + tool := inventory.NewServerToolWithContextHandlerAndSchemaOptions( |
| 272 | + mcp.Tool{ |
| 273 | + Name: probeToolName, |
| 274 | + InputSchema: &jsonschema.Schema{ |
| 275 | + Type: "object", |
| 276 | + Properties: map[string]*jsonschema.Schema{ |
| 277 | + "owner": {Type: "string"}, |
| 278 | + "repo": {Type: "string"}, |
| 279 | + "mode": {Type: "string", Enum: []any{"valid"}}, |
| 280 | + }, |
| 281 | + Required: []string{"owner", "repo", "mode"}, |
| 282 | + }, |
| 283 | + }, |
| 284 | + inventory.ToolsetMetadata{ID: "test"}, |
| 285 | + func(context.Context, *mcp.CallToolRequest, input) (*mcp.CallToolResult, oauthProbeOutput, error) { |
| 286 | + handlerCalls++ |
| 287 | + return nil, oauthProbeOutput{Status: "tool-ran"}, nil |
| 288 | + }, |
| 289 | + inventory.TypedSchemaOptions{ |
| 290 | + Preflight: func(ctx context.Context, _ *mcp.CallToolRequest) (context.Context, *mcp.CallToolResult, error) { |
| 291 | + preflightCalls++ |
| 292 | + return ctx, nil, nil |
| 293 | + }, |
| 294 | + }, |
| 295 | + func(raw json.RawMessage) (json.RawMessage, error) { |
| 296 | + normalizerCalls++ |
| 297 | + return raw, nil |
| 298 | + }, |
| 299 | + ) |
| 300 | + server := mcp.NewServer(&mcp.Implementation{Name: "oauth-http", Version: "test"}, nil) |
| 301 | + tool.RegisterFunc(server, nil, createOAuthToolMiddleware(fake, discardLogger())) |
| 302 | + server.AddReceivingMiddleware(func(next mcp.MethodHandler) mcp.MethodHandler { |
| 303 | + return func(ctx context.Context, method string, req mcp.Request) (mcp.Result, error) { |
| 304 | + if method == inventory.MCPMethodToolsCall { |
| 305 | + receivingCalls++ |
| 306 | + } |
| 307 | + return next(ctx, method, req) |
| 308 | + } |
| 309 | + }) |
| 310 | + handler := mcp.NewStreamableHTTPHandler(func(*http.Request) *mcp.Server { |
| 311 | + return server |
| 312 | + }, &mcp.StreamableHTTPOptions{Stateless: true, JSONResponse: true}) |
| 313 | + call := func(accept bool) *httptest.ResponseRecorder { |
| 314 | + t.Helper() |
| 315 | + params := map[string]any{ |
| 316 | + "name": probeToolName, "arguments": json.RawMessage(tc.arguments), |
| 317 | + "_meta": mcp.Meta{ |
| 318 | + mcp.MetaKeyProtocolVersion: inventory.ProtocolVersionMultiRoundTrip, |
| 319 | + mcp.MetaKeyClientInfo: &mcp.Implementation{Name: "test", Version: "test"}, |
| 320 | + mcp.MetaKeyClientCapabilities: &mcp.ClientCapabilities{Elicitation: &mcp.ElicitationCapabilities{URL: &mcp.URLElicitationCapabilities{}}}, |
| 321 | + }, |
| 322 | + } |
| 323 | + if accept { |
| 324 | + params["inputResponses"] = mcp.InputResponseMap{ |
| 325 | + oauthElicitIDPrefix + "flow-1": &mcp.ElicitResult{Action: "accept"}, |
| 326 | + } |
| 327 | + } |
| 328 | + body, err := json.Marshal(params) |
| 329 | + require.NoError(t, err) |
| 330 | + req := httptest.NewRequest(http.MethodPost, "/", bytes.NewBufferString(fmt.Sprintf( |
| 331 | + `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":%s}`, body))) |
| 332 | + req.Header.Set("Content-Type", "application/json") |
| 333 | + req.Header.Set("Accept", "application/json, text/event-stream") |
| 334 | + req.Header.Set("Mcp-Protocol-Version", inventory.ProtocolVersionMultiRoundTrip) |
| 335 | + req.Header.Set("Mcp-Method", "tools/call") |
| 336 | + req.Header.Set("Mcp-Name", probeToolName) |
| 337 | + req.Header.Set("Mcp-Param-owner", tc.owner) |
| 338 | + req.Header.Set("Mcp-Param-repo", tc.repo) |
| 339 | + rec := httptest.NewRecorder() |
| 340 | + handler.ServeHTTP(rec, req) |
| 341 | + return rec |
| 342 | + } |
| 343 | + rec := call(false) |
| 344 | + if tc.mismatch { |
| 345 | + require.Equal(t, http.StatusBadRequest, rec.Code, rec.Body.String()) |
| 346 | + assert.Contains(t, rec.Body.String(), "header mismatch") |
| 347 | + assert.Zero(t, receivingCalls) |
| 348 | + assert.Zero(t, fake.authCalls) |
| 349 | + } else { |
| 350 | + require.Equal(t, http.StatusOK, rec.Code, rec.Body.String()) |
| 351 | + assert.Contains(t, rec.Body.String(), `"resultType":"input_required"`) |
| 352 | + assert.NotContains(t, rec.Body.String(), "structuredContent") |
| 353 | + assert.Equal(t, 1, receivingCalls) |
| 354 | + assert.Equal(t, 1, fake.authCalls) |
| 355 | + } |
| 356 | + assert.Zero(t, handlerCalls) |
| 357 | + assert.Zero(t, preflightCalls) |
| 358 | + assert.Zero(t, normalizerCalls) |
| 359 | + if tc.mismatch { |
| 360 | + return |
| 361 | + } |
| 362 | + rec = call(true) |
| 363 | + require.Equal(t, http.StatusOK, rec.Code, rec.Body.String()) |
| 364 | + assert.Contains(t, rec.Body.String(), `"resultType":"complete"`) |
| 365 | + assert.NotContains(t, rec.Body.String(), `"resultType":"input_required"`) |
| 366 | + assert.Equal(t, 1, fake.awaitCalls) |
| 367 | + assert.Equal(t, 1, preflightCalls) |
| 368 | + assert.Equal(t, 1, normalizerCalls) |
| 369 | + if tc.valid { |
| 370 | + assert.Equal(t, 1, handlerCalls) |
| 371 | + assert.Contains(t, rec.Body.String(), `"structuredContent":{"status":"tool-ran"}`) |
| 372 | + } else { |
| 373 | + assert.Zero(t, handlerCalls) |
| 374 | + assert.Contains(t, rec.Body.String(), `"isError":true`) |
| 375 | + assert.NotContains(t, rec.Body.String(), "structuredContent") |
| 376 | + } |
| 377 | + }) |
| 378 | + } |
| 379 | +} |
| 380 | + |
240 | 381 | func TestOAuthTypedRegistrationGuardsInvalidArguments(t *testing.T) { |
241 | 382 | type input struct { |
242 | 383 | Mode string `json:"mode"` |
|
0 commit comments