| 1 | package plugin |
| 2 | |
| 3 | import ( |
| 4 | "context" |
| 5 | "encoding/json" |
| 6 | "errors" |
| 7 | "fmt" |
| 8 | "io" |
| 9 | "net/http" |
| 10 | "net/http/httptest" |
| 11 | "net/url" |
| 12 | "os" |
| 13 | "path/filepath" |
| 14 | "runtime" |
| 15 | "strings" |
| 16 | "sync" |
| 17 | "sync/atomic" |
| 18 | "testing" |
| 19 | "time" |
| 20 | ) |
| 21 | |
| 22 | func TestParseBearerChallenge(t *testing.T) { |
| 23 | metadata, scope, ok := parseBearerChallenge(`Basic realm="legacy", Bearer resource_metadata="https://mcp.example.test/.well-known/oauth-protected-resource", scope="mcp:connect files:read"`) |
| 24 | if !ok { |
| 25 | t.Fatal("Bearer challenge was not parsed") |
| 26 | } |
| 27 | if metadata != "https://mcp.example.test/.well-known/oauth-protected-resource" { |
| 28 | t.Fatalf("resource metadata = %q", metadata) |
| 29 | } |
| 30 | if scope != "mcp:connect files:read" { |
| 31 | t.Fatalf("scope = %q", scope) |
| 32 | } |
| 33 | } |
| 34 | |
| 35 | func TestPKCEChallengeMatchesRFC7636KnownAnswer(t *testing.T) { |
| 36 | const verifier = "dBjftJeZ4CVP-mB92K27uhbUJU1p1r_wW1gFWFOEjXk" |
| 37 | const want = "E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw-cM" |
| 38 | if got := pkceChallenge(verifier); got != want { |
| 39 | t.Fatalf("PKCE challenge = %q, want %q", got, want) |
| 40 | } |
| 41 | } |
| 42 | |
| 43 | func TestOAuthHTTPClientDoesNotChangeRuntimeIdentity(t *testing.T) { |
| 44 | base := Spec{Name: "remote", Type: "http", URL: "https://mcp.example.test"} |
| 45 | withClient := base |
| 46 | withClient.OAuthHTTPClient = &http.Client{} |
| 47 | if !MCPRuntimeSpecMatches(base, withClient) { |
| 48 | t.Fatal("host-local OAuth HTTP client changed MCP runtime identity") |
| 49 | } |
| 50 | } |
| 51 | |
| 52 | func TestAuthorizeHTTPMCPRejectsStaticAuthorizationHeader(t *testing.T) { |
| 53 | opened := false |
| 54 | err := AuthorizeHTTPMCP(context.Background(), Spec{ |
| 55 | Name: "remote", Type: "http", URL: "https://example.test/mcp", StateDir: t.TempDir(), |
| 56 | Headers: map[string]string{"Authorization": "Bearer configured"}, |
| 57 | }, func(string) error { |
| 58 | opened = true |
| 59 | return nil |
| 60 | }) |
| 61 | if err == nil || !strings.Contains(err.Error(), "explicit authentication") { |
| 62 | t.Fatalf("AuthorizeHTTPMCP error = %v", err) |
| 63 | } |
| 64 | if opened { |
| 65 | t.Fatal("static Authorization configuration opened the OAuth browser") |
| 66 | } |
| 67 | } |
| 68 | |
| 69 | func TestAuthorizeHTTPMCPRejectsStaticAPIKeyHeader(t *testing.T) { |
| 70 | opened := false |
| 71 | err := AuthorizeHTTPMCP(context.Background(), Spec{ |
| 72 | Name: "remote", Type: "http", URL: "https://example.test/mcp", StateDir: t.TempDir(), |
| 73 | Headers: map[string]string{"X-API-Key": "configured"}, |
| 74 | }, func(string) error { |
| 75 | opened = true |
| 76 | return nil |
| 77 | }) |
| 78 | if err == nil || !strings.Contains(err.Error(), "explicit authentication") { |
| 79 | t.Fatalf("AuthorizeHTTPMCP error = %v", err) |
| 80 | } |
| 81 | if opened { |
| 82 | t.Fatal("static API key configuration opened the OAuth browser") |
| 83 | } |
| 84 | } |
| 85 | |
| 86 | func TestAuthorizeHTTPMCPUsesDiscoveryPKCEAndPersistsPrivateToken(t *testing.T) { |
| 87 | stateDir := t.TempDir() |
| 88 | var server *httptest.Server |
| 89 | var mu sync.Mutex |
| 90 | registeredRedirect := "" |
| 91 | codeChallenge := "" |
| 92 | server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { |
| 93 | switch r.URL.Path { |
| 94 | case "/mcp": |
| 95 | if r.Header.Get("Authorization") != "Bearer access-one" { |
| 96 | w.Header().Set("WWW-Authenticate", fmt.Sprintf(`Bearer resource_metadata=%q, scope="mcp:connect"`, server.URL+"/.well-known/oauth-protected-resource")) |
| 97 | http.Error(w, "unauthorized", http.StatusUnauthorized) |
| 98 | return |
| 99 | } |
| 100 | writeOAuthMCPFixtureResponse(w, r) |
| 101 | case "/.well-known/oauth-protected-resource": |
| 102 | _ = json.NewEncoder(w).Encode(map[string]any{ |
| 103 | "resource": server.URL + "/mcp", |
| 104 | "authorization_servers": []string{server.URL}, |
| 105 | "scopes_supported": []string{"mcp:connect"}, |
| 106 | }) |
| 107 | case "/.well-known/oauth-authorization-server": |
| 108 | _ = json.NewEncoder(w).Encode(map[string]any{ |
| 109 | "issuer": server.URL, |
| 110 | "authorization_endpoint": server.URL + "/authorize", |
| 111 | "token_endpoint": server.URL + "/token", |
| 112 | "registration_endpoint": server.URL + "/register", |
| 113 | "code_challenge_methods_supported": []string{"S256"}, |
| 114 | "token_endpoint_auth_methods_supported": []string{"client_secret_basic"}, |
| 115 | }) |
| 116 | case "/register": |
| 117 | var registration map[string]any |
| 118 | if err := json.NewDecoder(r.Body).Decode(®istration); err != nil { |
| 119 | t.Errorf("decode registration: %v", err) |
| 120 | http.Error(w, "bad registration", http.StatusBadRequest) |
| 121 | return |
| 122 | } |
| 123 | redirects, _ := registration["redirect_uris"].([]any) |
| 124 | if len(redirects) != 1 { |
| 125 | t.Errorf("redirect_uris = %#v", registration["redirect_uris"]) |
| 126 | } else { |
| 127 | registeredRedirect, _ = redirects[0].(string) |
| 128 | } |
| 129 | _ = json.NewEncoder(w).Encode(map[string]any{ |
| 130 | "client_id": "reasonix-test", |
| 131 | "client_secret": "client-secret", |
| 132 | "token_endpoint_auth_method": "client_secret_basic", |
| 133 | }) |
| 134 | case "/token": |
| 135 | if user, pass, ok := r.BasicAuth(); !ok || user != "reasonix-test" || pass != "client-secret" { |
| 136 | t.Errorf("token endpoint client authentication = (%q, %q, %v)", user, pass, ok) |
| 137 | } |
| 138 | if err := r.ParseForm(); err != nil { |
| 139 | t.Errorf("parse token form: %v", err) |
| 140 | } |
| 141 | verifier := r.Form.Get("code_verifier") |
| 142 | mu.Lock() |
| 143 | expectedChallenge := codeChallenge |
| 144 | mu.Unlock() |
| 145 | if verifier == "" || pkceChallenge(verifier) != expectedChallenge { |
| 146 | t.Errorf("PKCE verifier does not match challenge") |
| 147 | } |
| 148 | if got := r.Form.Get("resource"); got != server.URL+"/mcp" { |
| 149 | t.Errorf("token resource = %q", got) |
| 150 | } |
| 151 | _ = json.NewEncoder(w).Encode(map[string]any{ |
| 152 | "access_token": "access-one", |
| 153 | "refresh_token": "refresh-one", |
| 154 | "token_type": "Bearer", |
| 155 | "expires_in": 3600, |
| 156 | "scope": "mcp:connect", |
| 157 | }) |
| 158 | default: |
| 159 | http.NotFound(w, r) |
| 160 | } |
| 161 | })) |
| 162 | defer server.Close() |
| 163 | |
| 164 | var oauthRequests atomic.Int32 |
| 165 | spec := Spec{ |
| 166 | Name: "figma", Type: "http", URL: server.URL + "/mcp", StateDir: stateDir, |
| 167 | OAuthHTTPClient: &http.Client{Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) { |
| 168 | oauthRequests.Add(1) |
| 169 | return http.DefaultTransport.RoundTrip(req) |
| 170 | })}, |
| 171 | } |
| 172 | openURL := func(raw string) error { |
| 173 | authURL, err := url.Parse(raw) |
| 174 | if err != nil { |
| 175 | return err |
| 176 | } |
| 177 | if authURL.Path != "/authorize" { |
| 178 | return fmt.Errorf("authorization path = %q", authURL.Path) |
| 179 | } |
| 180 | query := authURL.Query() |
| 181 | if query.Get("code_challenge_method") != "S256" { |
| 182 | return fmt.Errorf("code challenge method = %q", query.Get("code_challenge_method")) |
| 183 | } |
| 184 | if query.Get("resource") != server.URL+"/mcp" { |
| 185 | return fmt.Errorf("authorization resource = %q", query.Get("resource")) |
| 186 | } |
| 187 | mu.Lock() |
| 188 | codeChallenge = query.Get("code_challenge") |
| 189 | mu.Unlock() |
| 190 | callback, err := url.Parse(query.Get("redirect_uri")) |
| 191 | if err != nil { |
| 192 | return err |
| 193 | } |
| 194 | values := callback.Query() |
| 195 | values.Set("code", "authorization-code") |
| 196 | values.Set("state", query.Get("state")) |
| 197 | callback.RawQuery = values.Encode() |
| 198 | go func() { |
| 199 | resp, err := http.Get(callback.String()) |
| 200 | if err == nil { |
| 201 | _ = resp.Body.Close() |
| 202 | } |
| 203 | }() |
| 204 | return nil |
| 205 | } |
| 206 | |
| 207 | ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) |
| 208 | defer cancel() |
| 209 | if err := AuthorizeHTTPMCP(ctx, spec, openURL); err != nil { |
| 210 | t.Fatalf("AuthorizeHTTPMCP: %v", err) |
| 211 | } |
| 212 | if oauthRequests.Load() == 0 { |
| 213 | t.Fatal("AuthorizeHTTPMCP did not use the injected proxy-aware HTTP client") |
| 214 | } |
| 215 | if !strings.HasPrefix(registeredRedirect, "http://127.0.0.1:") { |
| 216 | t.Fatalf("registered redirect = %q", registeredRedirect) |
| 217 | } |
| 218 | tokenPath := filepath.Join(stateDir, mcpOAuthStateFile) |
| 219 | info, err := os.Stat(tokenPath) |
| 220 | if err != nil { |
| 221 | t.Fatalf("stat token state: %v", err) |
| 222 | } |
| 223 | // Windows has no Unix permission bits: os.WriteFile's 0600 intent is |
| 224 | // unobservable there (mode reports 0666), so the permission contract is |
| 225 | // asserted only where it exists. |
| 226 | if runtime.GOOS != "windows" { |
| 227 | if got := info.Mode().Perm(); got != 0o600 { |
| 228 | t.Fatalf("token state mode = %o, want 600", got) |
| 229 | } |
| 230 | } |
| 231 | |
| 232 | transport, err := newHTTPTransport(spec) |
| 233 | if err != nil { |
| 234 | t.Fatal(err) |
| 235 | } |
| 236 | defer transport.close() |
| 237 | result, err := transport.call(context.Background(), "ping", map[string]any{}) |
| 238 | if err != nil { |
| 239 | t.Fatalf("authenticated MCP call: %v", err) |
| 240 | } |
| 241 | if string(result) != `{}` { |
| 242 | t.Fatalf("result = %s, want typed empty ping result", result) |
| 243 | } |
| 244 | } |
| 245 | |
| 246 | func TestAuthorizeHTTPMCPDoesNotHoldStateLockDuringBrowser(t *testing.T) { |
| 247 | stateDir := t.TempDir() |
| 248 | const endpoint = "https://mcp.example.test/mcp" |
| 249 | client := &http.Client{Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) { |
| 250 | response := func(status int, body string) (*http.Response, error) { |
| 251 | return &http.Response{ |
| 252 | StatusCode: status, |
| 253 | Header: make(http.Header), |
| 254 | Body: io.NopCloser(strings.NewReader(body)), |
| 255 | Request: req, |
| 256 | }, nil |
| 257 | } |
| 258 | switch req.URL.Path { |
| 259 | case "/mcp": |
| 260 | resp, err := response(http.StatusUnauthorized, `unauthorized`) |
| 261 | if err != nil { |
| 262 | return nil, err |
| 263 | } |
| 264 | resp.Header.Set("WWW-Authenticate", `Bearer resource_metadata="https://mcp.example.test/.well-known/oauth-protected-resource"`) |
| 265 | return resp, nil |
| 266 | case "/.well-known/oauth-protected-resource": |
| 267 | return response(http.StatusOK, `{"resource":"https://mcp.example.test/mcp","authorization_servers":["https://mcp.example.test"],"scopes_supported":["mcp:connect"]}`) |
| 268 | case "/.well-known/oauth-authorization-server": |
| 269 | return response(http.StatusOK, `{"issuer":"https://mcp.example.test","authorization_endpoint":"https://mcp.example.test/authorize","token_endpoint":"https://mcp.example.test/token","registration_endpoint":"https://mcp.example.test/register","code_challenge_methods_supported":["S256"],"token_endpoint_auth_methods_supported":["client_secret_basic"]}`) |
| 270 | case "/register": |
| 271 | return response(http.StatusOK, `{"client_id":"reasonix-test","client_secret":"client-secret","token_endpoint_auth_method":"client_secret_basic"}`) |
| 272 | case "/token": |
| 273 | return response(http.StatusOK, `{"access_token":"access-one","refresh_token":"refresh-one","token_type":"Bearer","expires_in":3600}`) |
| 274 | default: |
| 275 | return response(http.StatusNotFound, `not found`) |
| 276 | } |
| 277 | })} |
| 278 | |
| 279 | openURL := func(raw string) error { |
| 280 | authURL, err := url.Parse(raw) |
| 281 | if err != nil { |
| 282 | return err |
| 283 | } |
| 284 | clearDone := make(chan error, 1) |
| 285 | go func() { |
| 286 | _, clearErr := ClearHTTPMCPOAuth(Spec{StateDir: stateDir}) |
| 287 | clearDone <- clearErr |
| 288 | }() |
| 289 | select { |
| 290 | case clearErr := <-clearDone: |
| 291 | if clearErr != nil { |
| 292 | return fmt.Errorf("clear during browser flow: %w", clearErr) |
| 293 | } |
| 294 | case <-time.After(time.Second): |
| 295 | return fmt.Errorf("clear during browser flow blocked on OAuth state lock") |
| 296 | } |
| 297 | callback, err := url.Parse(authURL.Query().Get("redirect_uri")) |
| 298 | if err != nil { |
| 299 | return err |
| 300 | } |
| 301 | query := callback.Query() |
| 302 | query.Set("code", "authorization-code") |
| 303 | query.Set("state", authURL.Query().Get("state")) |
| 304 | callback.RawQuery = query.Encode() |
| 305 | resp, err := http.Get(callback.String()) |
| 306 | if err == nil { |
| 307 | _ = resp.Body.Close() |
| 308 | } |
| 309 | return err |
| 310 | } |
| 311 | |
| 312 | ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) |
| 313 | defer cancel() |
| 314 | err := AuthorizeHTTPMCP(ctx, Spec{ |
| 315 | Name: "remote", Type: "http", URL: endpoint, StateDir: stateDir, |
| 316 | OAuthHTTPClient: client, |
| 317 | }, openURL) |
| 318 | if err == nil || !strings.Contains(err.Error(), "invalidated") { |
| 319 | t.Fatalf("AuthorizeHTTPMCP after concurrent clear = %v, want invalidation", err) |
| 320 | } |
| 321 | if _, err := os.Stat(mcpOAuthStatePath(stateDir)); !errors.Is(err, os.ErrNotExist) { |
| 322 | t.Fatalf("OAuth state was written after concurrent clear: %v", err) |
| 323 | } |
| 324 | } |
| 325 | |
| 326 | func TestHTTPMCPRefreshesExpiredTokenAndRotatesRefreshToken(t *testing.T) { |
| 327 | stateDir := t.TempDir() |
| 328 | refreshCalls := 0 |
| 329 | server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { |
| 330 | switch r.URL.Path { |
| 331 | case "/token": |
| 332 | refreshCalls++ |
| 333 | if err := r.ParseForm(); err != nil { |
| 334 | t.Fatal(err) |
| 335 | } |
| 336 | if r.Form.Get("grant_type") != "refresh_token" || r.Form.Get("refresh_token") != "refresh-old" { |
| 337 | t.Errorf("refresh form = %v", r.Form) |
| 338 | } |
| 339 | _ = json.NewEncoder(w).Encode(map[string]any{ |
| 340 | "access_token": "access-new", "refresh_token": "refresh-new", "token_type": "Bearer", "expires_in": 3600, |
| 341 | }) |
| 342 | case "/mcp": |
| 343 | if r.Header.Get("Authorization") != "Bearer access-new" { |
| 344 | http.Error(w, "unauthorized", http.StatusUnauthorized) |
| 345 | return |
| 346 | } |
| 347 | writeOAuthMCPFixtureResponse(w, r) |
| 348 | default: |
| 349 | http.NotFound(w, r) |
| 350 | } |
| 351 | })) |
| 352 | defer server.Close() |
| 353 | |
| 354 | state := mcpOAuthState{ |
| 355 | Version: 1, Resource: server.URL + "/mcp", Issuer: server.URL, |
| 356 | ClientID: "client", ClientSecret: "secret", TokenEndpoint: server.URL + "/token", TokenEndpointAuthMethod: "client_secret_basic", |
| 357 | AccessToken: "access-old", RefreshToken: "refresh-old", TokenType: "Bearer", Expiry: time.Now().Add(-time.Minute), |
| 358 | } |
| 359 | if err := saveMCPOAuthState(stateDir, state); err != nil { |
| 360 | t.Fatal(err) |
| 361 | } |
| 362 | transport, err := newHTTPTransport(Spec{Name: "remote", Type: "http", URL: server.URL + "/mcp", StateDir: stateDir}) |
| 363 | if err != nil { |
| 364 | t.Fatal(err) |
| 365 | } |
| 366 | defer transport.close() |
| 367 | if _, err := transport.call(context.Background(), "ping", nil); err != nil { |
| 368 | t.Fatalf("call after refresh: %v", err) |
| 369 | } |
| 370 | if refreshCalls != 1 { |
| 371 | t.Fatalf("refresh calls = %d, want 1", refreshCalls) |
| 372 | } |
| 373 | rotated, err := loadMCPOAuthState(stateDir) |
| 374 | if err != nil { |
| 375 | t.Fatal(err) |
| 376 | } |
| 377 | if rotated.RefreshToken != "refresh-new" || rotated.AccessToken != "access-new" { |
| 378 | t.Fatalf("rotated token state = %+v", rotated) |
| 379 | } |
| 380 | } |
| 381 | |
| 382 | func TestOAuthClientSecretBasicFormEncodesCredentials(t *testing.T) { |
| 383 | server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { |
| 384 | user, pass, ok := r.BasicAuth() |
| 385 | if !ok || user != "client+id%2B" || pass != "secret%3Avalue%2Fwith+space" { |
| 386 | t.Errorf("OAuth Basic credentials = (%q, %q, %v)", user, pass, ok) |
| 387 | } |
| 388 | _ = json.NewEncoder(w).Encode(map[string]any{ |
| 389 | "access_token": "access-new", "token_type": "Bearer", "expires_in": 3600, |
| 390 | }) |
| 391 | })) |
| 392 | defer server.Close() |
| 393 | |
| 394 | _, err := requestOAuthToken(context.Background(), server.Client(), mcpOAuthState{ |
| 395 | TokenEndpoint: server.URL, ClientID: "client id+", ClientSecret: "secret:value/with space", TokenEndpointAuthMethod: "client_secret_basic", |
| 396 | }, url.Values{"grant_type": {"authorization_code"}}) |
| 397 | if err != nil { |
| 398 | t.Fatalf("requestOAuthToken: %v", err) |
| 399 | } |
| 400 | } |
| 401 | |
| 402 | func TestHTTPMCPSerializesSharedRefreshTokenRotation(t *testing.T) { |
| 403 | stateDir := t.TempDir() |
| 404 | refreshStarted := make(chan struct{}) |
| 405 | allowRefresh := make(chan struct{}) |
| 406 | var refreshCalls atomic.Int32 |
| 407 | server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { |
| 408 | switch r.URL.Path { |
| 409 | case "/token": |
| 410 | call := refreshCalls.Add(1) |
| 411 | if call != 1 { |
| 412 | t.Errorf("refresh endpoint called %d times", call) |
| 413 | http.Error(w, "duplicate refresh", http.StatusBadRequest) |
| 414 | return |
| 415 | } |
| 416 | close(refreshStarted) |
| 417 | <-allowRefresh |
| 418 | if err := r.ParseForm(); err != nil { |
| 419 | t.Errorf("parse refresh form: %v", err) |
| 420 | } |
| 421 | if got := r.Form.Get("refresh_token"); got != "refresh-old" { |
| 422 | t.Errorf("refresh token = %q, want refresh-old", got) |
| 423 | } |
| 424 | _ = json.NewEncoder(w).Encode(map[string]any{ |
| 425 | "access_token": "access-new", "refresh_token": "refresh-new", "token_type": "Bearer", "expires_in": 3600, |
| 426 | }) |
| 427 | case "/mcp": |
| 428 | if r.Header.Get("Authorization") != "Bearer access-new" { |
| 429 | http.Error(w, "unauthorized", http.StatusUnauthorized) |
| 430 | return |
| 431 | } |
| 432 | writeOAuthMCPFixtureResponse(w, r) |
| 433 | default: |
| 434 | http.NotFound(w, r) |
| 435 | } |
| 436 | })) |
| 437 | defer server.Close() |
| 438 | |
| 439 | if err := saveMCPOAuthState(stateDir, mcpOAuthState{ |
| 440 | Version: 1, Resource: server.URL + "/mcp", Issuer: server.URL, |
| 441 | ClientID: "client", ClientSecret: "secret", TokenEndpoint: server.URL + "/token", TokenEndpointAuthMethod: "client_secret_basic", |
| 442 | AccessToken: "access-old", RefreshToken: "refresh-old", TokenType: "Bearer", Expiry: time.Now().Add(-time.Minute), |
| 443 | }); err != nil { |
| 444 | t.Fatal(err) |
| 445 | } |
| 446 | spec := Spec{Name: "remote", Type: "http", URL: server.URL + "/mcp", StateDir: stateDir} |
| 447 | first, err := newHTTPTransport(spec) |
| 448 | if err != nil { |
| 449 | t.Fatal(err) |
| 450 | } |
| 451 | defer first.close() |
| 452 | second, err := newHTTPTransport(spec) |
| 453 | if err != nil { |
| 454 | t.Fatal(err) |
| 455 | } |
| 456 | defer second.close() |
| 457 | |
| 458 | errs := make(chan error, 2) |
| 459 | go func() { |
| 460 | _, err := first.call(context.Background(), "ping", nil) |
| 461 | errs <- err |
| 462 | }() |
| 463 | <-refreshStarted |
| 464 | secondStarted := make(chan struct{}) |
| 465 | go func() { |
| 466 | close(secondStarted) |
| 467 | _, err := second.call(context.Background(), "ping", nil) |
| 468 | errs <- err |
| 469 | }() |
| 470 | <-secondStarted |
| 471 | close(allowRefresh) |
| 472 | for range 2 { |
| 473 | if err := <-errs; err != nil { |
| 474 | t.Fatalf("shared refresh call: %v", err) |
| 475 | } |
| 476 | } |
| 477 | if got := refreshCalls.Load(); got != 1 { |
| 478 | t.Fatalf("refresh calls = %d, want 1", got) |
| 479 | } |
| 480 | } |
| 481 | |
| 482 | func TestMCPOAuthConcurrentUnauthorizedRefreshesUnexpiredTokenOnce(t *testing.T) { |
| 483 | stateDir := t.TempDir() |
| 484 | refreshStarted := make(chan struct{}) |
| 485 | allowRefresh := make(chan struct{}) |
| 486 | var refreshCalls atomic.Int32 |
| 487 | server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { |
| 488 | if r.URL.Path != "/token" { |
| 489 | http.NotFound(w, r) |
| 490 | return |
| 491 | } |
| 492 | if call := refreshCalls.Add(1); call != 1 { |
| 493 | t.Errorf("refresh endpoint called %d times", call) |
| 494 | } |
| 495 | if refreshCalls.Load() == 1 { |
| 496 | close(refreshStarted) |
| 497 | <-allowRefresh |
| 498 | } |
| 499 | _ = json.NewEncoder(w).Encode(map[string]any{ |
| 500 | "access_token": "access-new", "refresh_token": "refresh-new", "token_type": "Bearer", "expires_in": 3600, |
| 501 | }) |
| 502 | })) |
| 503 | defer server.Close() |
| 504 | |
| 505 | if err := saveMCPOAuthState(stateDir, mcpOAuthState{ |
| 506 | Version: 1, Resource: server.URL + "/mcp", Issuer: server.URL, |
| 507 | ClientID: "client", TokenEndpoint: server.URL + "/token", |
| 508 | AccessToken: "access-revoked", RefreshToken: "refresh-old", TokenType: "Bearer", Expiry: time.Now().Add(time.Hour), |
| 509 | }); err != nil { |
| 510 | t.Fatal(err) |
| 511 | } |
| 512 | client, err := newMCPOAuthClient(stateDir, server.Client()) |
| 513 | if err != nil { |
| 514 | t.Fatal(err) |
| 515 | } |
| 516 | |
| 517 | authorize := func() error { |
| 518 | request := httptest.NewRequest(http.MethodPost, server.URL+"/mcp", nil) |
| 519 | request.Header.Set("Authorization", "Bearer access-revoked") |
| 520 | response := &http.Response{StatusCode: http.StatusUnauthorized, Body: io.NopCloser(strings.NewReader("unauthorized"))} |
| 521 | return client.Authorize(t.Context(), request, response) |
| 522 | } |
| 523 | errs := make(chan error, 2) |
| 524 | go func() { errs <- authorize() }() |
| 525 | <-refreshStarted |
| 526 | go func() { errs <- authorize() }() |
| 527 | close(allowRefresh) |
| 528 | for range 2 { |
| 529 | if err := <-errs; err != nil { |
| 530 | t.Fatalf("Authorize: %v", err) |
| 531 | } |
| 532 | } |
| 533 | if got := refreshCalls.Load(); got != 1 { |
| 534 | t.Fatalf("refresh calls = %d, want one shared refresh", got) |
| 535 | } |
| 536 | if client.state.AccessToken != "access-new" { |
| 537 | t.Fatalf("OAuth client kept stale access token %q", client.state.AccessToken) |
| 538 | } |
| 539 | } |
| 540 | |
| 541 | func TestHTTPMCPRefreshReleasesCrossProcessLockDuringTokenRequest(t *testing.T) { |
| 542 | stateDir := t.TempDir() |
| 543 | refreshStarted := make(chan struct{}) |
| 544 | allowRefresh := make(chan struct{}) |
| 545 | server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { |
| 546 | if r.URL.Path != "/token" { |
| 547 | http.NotFound(w, r) |
| 548 | return |
| 549 | } |
| 550 | close(refreshStarted) |
| 551 | <-allowRefresh |
| 552 | _ = json.NewEncoder(w).Encode(map[string]any{ |
| 553 | "access_token": "access-new", "refresh_token": "refresh-new", "token_type": "Bearer", "expires_in": 3600, |
| 554 | }) |
| 555 | })) |
| 556 | defer server.Close() |
| 557 | if err := saveMCPOAuthState(stateDir, mcpOAuthState{ |
| 558 | Version: 1, Resource: server.URL + "/mcp", TokenEndpoint: server.URL + "/token", ClientID: "client", |
| 559 | AccessToken: "access-old", RefreshToken: "refresh-old", TokenType: "Bearer", Expiry: time.Now().Add(-time.Minute), |
| 560 | }); err != nil { |
| 561 | t.Fatal(err) |
| 562 | } |
| 563 | transport, err := newHTTPTransport(Spec{Name: "remote", Type: "http", URL: server.URL + "/mcp", StateDir: stateDir}) |
| 564 | if err != nil { |
| 565 | t.Fatal(err) |
| 566 | } |
| 567 | defer transport.close() |
| 568 | callDone := make(chan error, 1) |
| 569 | go func() { |
| 570 | _, callErr := transport.call(context.Background(), "ping", nil) |
| 571 | callDone <- callErr |
| 572 | }() |
| 573 | <-refreshStarted |
| 574 | |
| 575 | clearDone := make(chan error, 1) |
| 576 | go func() { |
| 577 | _, clearErr := ClearHTTPMCPOAuth(Spec{StateDir: stateDir}) |
| 578 | clearDone <- clearErr |
| 579 | }() |
| 580 | select { |
| 581 | case clearErr := <-clearDone: |
| 582 | if clearErr != nil { |
| 583 | t.Fatalf("ClearHTTPMCPOAuth during refresh: %v", clearErr) |
| 584 | } |
| 585 | case <-time.After(time.Second): |
| 586 | t.Fatal("ClearHTTPMCPOAuth blocked on the token endpoint") |
| 587 | } |
| 588 | close(allowRefresh) |
| 589 | if err := <-callDone; err == nil || !strings.Contains(err.Error(), "invalidated") { |
| 590 | t.Fatalf("refresh after clear error = %v, want invalidation", err) |
| 591 | } |
| 592 | if _, err := os.Stat(mcpOAuthStatePath(stateDir)); !errors.Is(err, os.ErrNotExist) { |
| 593 | t.Fatalf("cleared OAuth state was recreated: %v", err) |
| 594 | } |
| 595 | } |
| 596 | |
| 597 | func TestHTTPMCPRejectsOAuthStateForDifferentResource(t *testing.T) { |
| 598 | stateDir := t.TempDir() |
| 599 | if err := saveMCPOAuthState(stateDir, mcpOAuthState{ |
| 600 | Version: 1, Resource: "https://old.example.test/mcp", Issuer: "https://auth.example.test", |
| 601 | ClientID: "client", AccessToken: "must-not-leak", TokenType: "Bearer", |
| 602 | }); err != nil { |
| 603 | t.Fatal(err) |
| 604 | } |
| 605 | |
| 606 | _, err := newHTTPTransport(Spec{Name: "remote", Type: "http", URL: "https://new.example.test/mcp", StateDir: stateDir}) |
| 607 | if err == nil || !strings.Contains(err.Error(), "different MCP resource") { |
| 608 | t.Fatalf("newHTTPTransport error = %v, want resource-binding rejection", err) |
| 609 | } |
| 610 | } |
| 611 | |
| 612 | func TestSameCanonicalResourceRejectsURLUserinfo(t *testing.T) { |
| 613 | if sameCanonicalResource("https://user:pass@mcp.example.test/mcp", "https://mcp.example.test/mcp") { |
| 614 | t.Fatal("credentialed URL must not match an OAuth resource") |
| 615 | } |
| 616 | } |
| 617 | |
| 618 | func TestClearHTTPMCPOAuthRemovesOnlyReasonixState(t *testing.T) { |
| 619 | stateDir := t.TempDir() |
| 620 | if err := saveMCPOAuthState(stateDir, mcpOAuthState{ |
| 621 | Version: 1, Resource: "https://mcp.example.test/mcp", Issuer: "https://auth.example.test", |
| 622 | ClientID: "client", AccessToken: "access-token", TokenType: "Bearer", |
| 623 | }); err != nil { |
| 624 | t.Fatal(err) |
| 625 | } |
| 626 | neighbor := filepath.Join(stateDir, "session.json") |
| 627 | if err := os.WriteFile(neighbor, []byte("keep"), 0o600); err != nil { |
| 628 | t.Fatal(err) |
| 629 | } |
| 630 | |
| 631 | changed, err := ClearHTTPMCPOAuth(Spec{StateDir: stateDir}) |
| 632 | if err != nil { |
| 633 | t.Fatalf("ClearHTTPMCPOAuth: %v", err) |
| 634 | } |
| 635 | if !changed { |
| 636 | t.Fatal("ClearHTTPMCPOAuth reported no change") |
| 637 | } |
| 638 | if _, err := os.Stat(filepath.Join(stateDir, mcpOAuthStateFile)); !errors.Is(err, os.ErrNotExist) { |
| 639 | t.Fatalf("OAuth state still exists or stat failed: %v", err) |
| 640 | } |
| 641 | if got, err := os.ReadFile(neighbor); err != nil || string(got) != "keep" { |
| 642 | t.Fatalf("neighboring MCP state changed: data=%q err=%v", got, err) |
| 643 | } |
| 644 | changed, err = ClearHTTPMCPOAuth(Spec{StateDir: stateDir}) |
| 645 | if err != nil || changed { |
| 646 | t.Fatalf("second ClearHTTPMCPOAuth = (%v, %v), want (false, nil)", changed, err) |
| 647 | } |
| 648 | } |
| 649 | |
| 650 | func TestClearHTTPMCPOAuthAllowsMissingPrivateStateDirectory(t *testing.T) { |
| 651 | stateDir := filepath.Join(t.TempDir(), "not-created-yet") |
| 652 | changed, err := ClearHTTPMCPOAuth(Spec{StateDir: stateDir}) |
| 653 | if err != nil || changed { |
| 654 | t.Fatalf("ClearHTTPMCPOAuth = (%v, %v), want (false, nil)", changed, err) |
| 655 | } |
| 656 | } |
| 657 | |
| 658 | func TestReconcileHTTPMCPOAuthAfterRemovalPreservesOnlyMatchingFallback(t *testing.T) { |
| 659 | stateDir := t.TempDir() |
| 660 | const resource = "https://mcp.example.test/mcp?workspace=main" |
| 661 | writeState := func() { |
| 662 | t.Helper() |
| 663 | if err := saveMCPOAuthState(stateDir, mcpOAuthState{ |
| 664 | Version: 1, Resource: resource, ClientID: "client", AccessToken: "access", TokenType: "Bearer", |
| 665 | }); err != nil { |
| 666 | t.Fatal(err) |
| 667 | } |
| 668 | } |
| 669 | writeState() |
| 670 | changed, err := ReconcileHTTPMCPOAuthAfterRemoval(Spec{StateDir: stateDir}, resource) |
| 671 | if err != nil || changed { |
| 672 | t.Fatalf("matching fallback reconciliation = (%v, %v), want (false, nil)", changed, err) |
| 673 | } |
| 674 | if _, err := os.Stat(mcpOAuthStatePath(stateDir)); err != nil { |
| 675 | t.Fatalf("matching fallback OAuth state was removed: %v", err) |
| 676 | } |
| 677 | |
| 678 | changed, err = ReconcileHTTPMCPOAuthAfterRemoval(Spec{StateDir: stateDir}, "https://other.example.test/mcp") |
| 679 | if err != nil || !changed { |
| 680 | t.Fatalf("different fallback reconciliation = (%v, %v), want (true, nil)", changed, err) |
| 681 | } |
| 682 | if _, err := os.Stat(mcpOAuthStatePath(stateDir)); !errors.Is(err, os.ErrNotExist) { |
| 683 | t.Fatalf("different fallback OAuth state still exists: %v", err) |
| 684 | } |
| 685 | } |
| 686 | |
| 687 | func TestMCPAuthGenerationInvalidatesPendingAuthorization(t *testing.T) { |
| 688 | stateDir := t.TempDir() |
| 689 | generation, err := captureMCPOAuthGeneration(context.Background(), stateDir) |
| 690 | if err != nil { |
| 691 | t.Fatalf("captureMCPOAuthGeneration: %v", err) |
| 692 | } |
| 693 | if err := bumpMCPOAuthGeneration(stateDir); err != nil { |
| 694 | t.Fatalf("bumpMCPOAuthGeneration: %v", err) |
| 695 | } |
| 696 | err = saveMCPOAuthStateIfGenerationUnchanged(context.Background(), stateDir, generation, mcpOAuthState{ |
| 697 | Resource: "https://mcp.example.test/mcp", AccessToken: "must-not-save", |
| 698 | }) |
| 699 | if err == nil || !strings.Contains(err.Error(), "invalidated") { |
| 700 | t.Fatalf("save after invalidation error = %v", err) |
| 701 | } |
| 702 | if _, err := os.Stat(mcpOAuthStatePath(stateDir)); !errors.Is(err, os.ErrNotExist) { |
| 703 | t.Fatalf("invalidated authorization wrote OAuth state: %v", err) |
| 704 | } |
| 705 | } |
| 706 | |
| 707 | func TestReconcileDifferentFallbackInvalidatesPendingAuthorizationWithoutState(t *testing.T) { |
| 708 | stateDir := t.TempDir() |
| 709 | generation, err := captureMCPOAuthGeneration(context.Background(), stateDir) |
| 710 | if err != nil { |
| 711 | t.Fatalf("captureMCPOAuthGeneration: %v", err) |
| 712 | } |
| 713 | changed, err := ReconcileHTTPMCPOAuthAfterRemoval(Spec{StateDir: stateDir}, "https://other.example.test/mcp") |
| 714 | if err != nil || changed { |
| 715 | t.Fatalf("reconcile without OAuth state = (%v, %v), want (false, nil)", changed, err) |
| 716 | } |
| 717 | err = saveMCPOAuthStateIfGenerationUnchanged(context.Background(), stateDir, generation, mcpOAuthState{ |
| 718 | Resource: "https://removed.example.test/mcp", AccessToken: "must-not-save", |
| 719 | }) |
| 720 | if err == nil || !strings.Contains(err.Error(), "invalidated") { |
| 721 | t.Fatalf("save after different fallback reconciliation error = %v", err) |
| 722 | } |
| 723 | } |
| 724 | |
| 725 | func TestClearedOAuthStateCannotBeResurrectedByStaleTransport(t *testing.T) { |
| 726 | stateDir := t.TempDir() |
| 727 | server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { |
| 728 | http.Error(w, "unauthorized", http.StatusUnauthorized) |
| 729 | })) |
| 730 | defer server.Close() |
| 731 | if err := saveMCPOAuthState(stateDir, mcpOAuthState{ |
| 732 | Version: 1, Resource: server.URL, Issuer: server.URL, TokenEndpoint: server.URL + "/token", |
| 733 | ClientID: "client", AccessToken: "access-old", RefreshToken: "refresh-old", TokenType: "Bearer", |
| 734 | }); err != nil { |
| 735 | t.Fatal(err) |
| 736 | } |
| 737 | transport, err := newHTTPTransport(Spec{Name: "remote", Type: "http", URL: server.URL, StateDir: stateDir}) |
| 738 | if err != nil { |
| 739 | t.Fatal(err) |
| 740 | } |
| 741 | defer transport.close() |
| 742 | if changed, err := ClearHTTPMCPOAuth(Spec{StateDir: stateDir}); err != nil || !changed { |
| 743 | t.Fatalf("ClearHTTPMCPOAuth = (%v, %v), want (true, nil)", changed, err) |
| 744 | } |
| 745 | if _, err := transport.call(context.Background(), "ping", nil); err == nil { |
| 746 | t.Fatal("stale transport call unexpectedly succeeded after clearing OAuth state") |
| 747 | } |
| 748 | if _, err := os.Stat(filepath.Join(stateDir, mcpOAuthStateFile)); !errors.Is(err, os.ErrNotExist) { |
| 749 | t.Fatalf("stale transport recreated OAuth state: %v", err) |
| 750 | } |
| 751 | } |
| 752 | |
| 753 | func TestOAuthErrorsRedactCredentialMaterial(t *testing.T) { |
| 754 | const secret = "fixture-oauth-secret-do-not-log-123456" |
| 755 | |
| 756 | resp := &http.Response{ |
| 757 | StatusCode: http.StatusBadRequest, |
| 758 | Body: io.NopCloser(strings.NewReader(`{"error":"invalid_token","access_token":"` + secret + `"}`)), |
| 759 | } |
| 760 | if got := oauthHTTPError("token request", resp).Error(); strings.Contains(got, secret) { |
| 761 | t.Fatalf("HTTP error leaked credential: %s", got) |
| 762 | } |
| 763 | |
| 764 | result := make(chan oauthCallbackResult, 1) |
| 765 | handler := oauthCallbackHandler("expected", result) |
| 766 | req := httptest.NewRequest(http.MethodGet, "/oauth/callback?state=expected&error=access_denied&error_description=token%3A"+secret, nil) |
| 767 | handler.ServeHTTP(httptest.NewRecorder(), req) |
| 768 | if callback := <-result; callback.Err == nil || strings.Contains(callback.Err.Error(), secret) { |
| 769 | t.Fatalf("callback error was not safely redacted: %v", callback.Err) |
| 770 | } |
| 771 | } |
| 772 | |
| 773 | type roundTripFunc func(*http.Request) (*http.Response, error) |
| 774 | |
| 775 | func (f roundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) { |
| 776 | return f(req) |
| 777 | } |
| 778 |