返回 DeepSeek-Reasonix
oauth_test.go
根目录 / internal / plugin / oauth_test.go
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(&registration); 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
778 lines GO