From dea22953f0937ce4054ad4aedc9735954c6d4003 Mon Sep 17 00:00:00 2001 From: PierrunoYT Date: Fri, 2 Oct 2026 19:48:44 +0200 Subject: [PATCH] fix(mcp): refresh OAuth tokens for discovered endpoints and registered clients Login discovers the token endpoint (RFC 9728 / 8414) and may obtain a client_id/client_secret through dynamic registration, but persisted none of it. Refresh read only the config-file OAuthConfig, which is empty for the standard `auth: oauth` setup, so every expired access token failed with "no token endpoint configured for refresh" and required a new `zero mcp oauth login`. Store the discovered token endpoint (and whether it was server-advertised) and the dynamically registered client on the token, and fall back to them in Refresh. Explicit config still wins, and values already in config are not copied into the store. A stored endpoint is re-validated before use: the shared https/loopback rule always, and for a server-advertised endpoint the same public-network policy and pinned client Login used. The refreshed token keeps the stored fields so later refreshes keep working. Fixes #1101 Co-Authored-By: Claude Sonnet 5.5 --- internal/mcp/network_client.go | 6 +- internal/mcp/oauth.go | 64 ++++- internal/mcp/oauth_refresh_stored_test.go | 315 ++++++++++++++++++++++ internal/mcp/oauth_store.go | 14 + internal/oauth/oauth.go | 10 + 5 files changed, 407 insertions(+), 2 deletions(-) create mode 100644 internal/mcp/oauth_refresh_stored_test.go diff --git a/internal/mcp/network_client.go b/internal/mcp/network_client.go index ce0b41162..a75d64ec1 100644 --- a/internal/mcp/network_client.go +++ b/internal/mcp/network_client.go @@ -871,7 +871,11 @@ func (source *storeTokenSource) Refresh(ctx context.Context) (string, error) { if !ok { return "", fmt.Errorf("no stored OAuth token for MCP server %s", source.server.Name) } - refreshed, err := refreshAccessToken(ctx, source.httpClient, source.config(), token, source.now) + cfg, client, err := refreshSettings(ctx, source.httpClient, source.server.URL, source.config(), token) + if err != nil { + return "", err + } + refreshed, err := refreshAccessToken(ctx, client, cfg, token, source.now) if err != nil { return "", err } diff --git a/internal/mcp/oauth.go b/internal/mcp/oauth.go index 05a3b296e..6eddc2db7 100644 --- a/internal/mcp/oauth.go +++ b/internal/mcp/oauth.go @@ -59,6 +59,11 @@ func tokenToStored(t oauth.Token) StoredToken { TokenType: t.TokenType, Scopes: t.Scopes, ExpiresAt: t.ExpiresAt, + + TokenEndpoint: t.TokenEndpoint, + ProtectedTokenEndpoint: t.ProtectedTokenEndpoint, + ClientID: t.ClientID, + ClientSecret: t.ClientSecret, } } @@ -341,7 +346,53 @@ func refreshAccessToken(ctx context.Context, client *http.Client, cfg OAuthConfi if err != nil { return StoredToken{}, err } - return tokenToStored(token), nil + refreshed := tokenToStored(token) + // The token response never carries what Login learned about the client and + // endpoint, so carry it over or the next refresh would lose it. + refreshed.TokenEndpoint = current.TokenEndpoint + refreshed.ProtectedTokenEndpoint = current.ProtectedTokenEndpoint + refreshed.ClientID = current.ClientID + refreshed.ClientSecret = current.ClientSecret + return refreshed, nil +} + +// refreshSettings fills the gaps in the configured OAuth settings from what the +// login stored: the discovered token endpoint and the dynamically registered +// client. Explicit config always wins. A stored endpoint is re-validated before +// credentials are posted to it, and one that came from server-advertised +// discovery is also held to the public-network policy, with the returned client +// pinned to it, exactly as Login did. +func refreshSettings(ctx context.Context, client *http.Client, resourceURL string, cfg OAuthConfig, stored StoredToken) (OAuthConfig, *http.Client, error) { + if strings.TrimSpace(cfg.ClientID) == "" && strings.TrimSpace(stored.ClientID) != "" { + cfg.ClientID = stored.ClientID + if strings.TrimSpace(cfg.ClientSecret) == "" { + cfg.ClientSecret = stored.ClientSecret + } + } + if strings.TrimSpace(cfg.TokenEndpoint) != "" || strings.TrimSpace(stored.TokenEndpoint) == "" { + return cfg, client, nil + } + endpoint := strings.TrimSpace(stored.TokenEndpoint) + if err := oauth.ValidateEndpointURL(endpoint); err != nil { + return cfg, client, fmt.Errorf("mcp oauth: stored token endpoint: %w", err) + } + if stored.ProtectedTokenEndpoint { + metadata := authServerMetadata{ + TokenEndpoint: endpoint, + protectedResourceURL: strings.TrimSpace(resourceURL), + protectTokenEndpoint: true, + } + if err := validateProtectedDiscoveredEndpoints(ctx, metadata); err != nil { + return cfg, client, err + } + protected, err := newAdvertisedOAuthDiscoveryClient(client, metadata.protectedResourceURL) + if err != nil { + return cfg, client, err + } + client = protected + } + cfg.TokenEndpoint = endpoint + return cfg, client, nil } // registerClient performs dynamic client registration against the registration @@ -446,6 +497,7 @@ func Login(ctx context.Context, options LoginOptions) (StoredToken, error) { tokenClient = protectedEndpointClient } + registeredClientID, registeredClientSecret := "", "" if strings.TrimSpace(cfg.ClientID) == "" { if registration := strings.TrimSpace(metadata.RegistrationEndpoint); registration != "" { registrationClient := httpClient @@ -457,8 +509,10 @@ func Login(ctx context.Context, options LoginOptions) (StoredToken, error) { return StoredToken{}, regErr } cfg.ClientID = clientID + registeredClientID = clientID if clientSecret != "" { cfg.ClientSecret = clientSecret + registeredClientSecret = clientSecret } } } @@ -537,6 +591,14 @@ func Login(ctx context.Context, options LoginOptions) (StoredToken, error) { if err != nil { return StoredToken{}, err } + // Persist only what config does not already carry, so refresh can reuse + // it without copying configured secrets into the token store. + if strings.TrimSpace(options.Config.TokenEndpoint) == "" { + token.TokenEndpoint = metadata.TokenEndpoint + token.ProtectedTokenEndpoint = metadata.protectTokenEndpoint + } + token.ClientID = registeredClientID + token.ClientSecret = registeredClientSecret return token, nil case <-loginCtx.Done(): return StoredToken{}, fmt.Errorf("timed out waiting for OAuth authorization callback: %w", loginCtx.Err()) diff --git a/internal/mcp/oauth_refresh_stored_test.go b/internal/mcp/oauth_refresh_stored_test.go new file mode 100644 index 000000000..4dcd1ed25 --- /dev/null +++ b/internal/mcp/oauth_refresh_stored_test.go @@ -0,0 +1,315 @@ +package mcp + +import ( + "context" + "encoding/json" + "errors" + "net/http" + "net/http/httptest" + "net/url" + "path/filepath" + "strings" + "sync/atomic" + "testing" + "time" + + "github.com/Gitlawb/zero/internal/oauth" +) + +const testServerIdentity = "0123456789abcdef0123456789abcdef" + +// discoveryOAuthServer serves a standard `auth: oauth` setup: protected-resource +// metadata, authorization-server metadata with a registration endpoint, dynamic +// client registration, and a token endpoint that records refresh requests. +type discoveryOAuthServer struct { + *httptest.Server + refreshForms chan url.Values +} + +func newDiscoveryOAuthServer(t *testing.T) *discoveryOAuthServer { + t.Helper() + fixture := &discoveryOAuthServer{refreshForms: make(chan url.Values, 4)} + fixture.Server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + base := "http://" + r.Host + switch r.URL.Path { + case "/mcp": + w.Header().Set("WWW-Authenticate", `Bearer resource_metadata="`+base+`/resource-metadata"`) + w.WriteHeader(http.StatusUnauthorized) + case "/resource-metadata": + _ = json.NewEncoder(w).Encode(map[string]any{ + "resource": base + "/mcp", + "authorization_servers": []string{base}, + }) + case "/.well-known/oauth-authorization-server": + _ = json.NewEncoder(w).Encode(map[string]any{ + "issuer": base, + "authorization_endpoint": base + "/authorize", + "token_endpoint": base + "/token", + "registration_endpoint": base + "/register", + }) + case "/register": + _ = json.NewEncoder(w).Encode(map[string]any{ + "client_id": "registered-client", + "client_secret": "registered-secret", + }) + case "/token": + if err := r.ParseForm(); err != nil { + t.Errorf("parse token form: %v", err) + } + if r.Form.Get("grant_type") == "refresh_token" { + fixture.refreshForms <- r.Form + _ = json.NewEncoder(w).Encode(map[string]any{ + "access_token": "access-refreshed", + "refresh_token": "refresh-rotated", + "token_type": "Bearer", + "expires_in": 3600, + }) + return + } + _ = json.NewEncoder(w).Encode(map[string]any{ + "access_token": "access-initial", + "refresh_token": "refresh-initial", + "token_type": "Bearer", + "expires_in": 3600, + }) + default: + http.NotFound(w, r) + } + })) + t.Cleanup(fixture.Close) + return fixture +} + +func driveLoopbackCallback(authURL string) error { + parsed, err := url.Parse(authURL) + if err != nil { + return err + } + callbackURL := parsed.Query().Get("redirect_uri") + "?code=auth-code&state=" + url.QueryEscape(parsed.Query().Get("state")) + go func() { + for attempt := 0; attempt < 20; attempt++ { + response, requestErr := http.Get(callbackURL) + if requestErr == nil { + response.Body.Close() + return + } + time.Sleep(25 * time.Millisecond) + } + }() + return nil +} + +func TestLoginRecordsDiscoveredEndpointAndRegisteredClient(t *testing.T) { + fixture := newDiscoveryOAuthServer(t) + + token, err := Login(context.Background(), LoginOptions{ + ServerName: "notion-style", + ServerURL: fixture.URL + "/mcp", + Config: OAuthConfig{}, + HTTPClient: fixture.Client(), + OpenBrowser: driveLoopbackCallback, + Timeout: 5 * time.Second, + Now: time.Now, + }) + if err != nil { + t.Fatalf("Login() error = %v", err) + } + if token.TokenEndpoint != fixture.URL+"/token" || !token.ProtectedTokenEndpoint { + t.Fatalf("stored token endpoint = %q protected=%v, want the discovered protected endpoint", token.TokenEndpoint, token.ProtectedTokenEndpoint) + } + if token.ClientID != "registered-client" || token.ClientSecret != "registered-secret" { + t.Fatalf("stored client = %q/%q, want the dynamically registered client", token.ClientID, token.ClientSecret) + } +} + +func TestLoginDoesNotCopyConfiguredClientOrEndpointIntoStore(t *testing.T) { + fixture := newDiscoveryOAuthServer(t) + + token, err := Login(context.Background(), LoginOptions{ + ServerName: "configured", + ServerURL: fixture.URL + "/mcp", + Config: OAuthConfig{ + ClientID: "configured-client", + ClientSecret: "configured-secret", + TokenEndpoint: fixture.URL + "/token", + }, + HTTPClient: fixture.Client(), + OpenBrowser: driveLoopbackCallback, + Timeout: 5 * time.Second, + Now: time.Now, + }) + if err != nil { + t.Fatalf("Login() error = %v", err) + } + if token.TokenEndpoint != "" || token.ProtectedTokenEndpoint || token.ClientID != "" || token.ClientSecret != "" { + t.Fatalf("configured values were copied into the token store: %#v", token) + } +} + +// The standard `auth: oauth` setup configures no token endpoint or client, so +// refresh must work from what the discovery-based login stored. +func TestStoreTokenSourceRefreshUsesStoredDiscoveryAndClient(t *testing.T) { + fixture := newDiscoveryOAuthServer(t) + server := Server{Name: "notion-style", URL: fixture.URL + "/mcp", Identity: testServerIdentity, Auth: ServerAuthOAuth} + + loggedIn, err := Login(context.Background(), LoginOptions{ + ServerName: server.Name, + ServerURL: server.URL, + HTTPClient: fixture.Client(), + OpenBrowser: driveLoopbackCallback, + Timeout: 5 * time.Second, + Now: time.Now, + }) + if err != nil { + t.Fatalf("Login() error = %v", err) + } + store, err := NewTokenStore(TokenStoreOptions{FilePath: filepath.Join(t.TempDir(), "oauth-tokens.json")}) + if err != nil { + t.Fatalf("NewTokenStore() error = %v", err) + } + if err := store.SaveForServer(server, loggedIn); err != nil { + t.Fatalf("SaveForServer() error = %v", err) + } + + source := &storeTokenSource{server: server, store: store, httpClient: http.DefaultClient, now: time.Now} + access, err := source.Refresh(context.Background()) + if err != nil { + t.Fatalf("Refresh() error = %v", err) + } + if access != "access-refreshed" { + t.Fatalf("Refresh() access token = %q", access) + } + form := <-fixture.refreshForms + if form.Get("client_id") != "registered-client" || form.Get("client_secret") != "registered-secret" || form.Get("refresh_token") != "refresh-initial" { + t.Fatalf("refresh form = %v, want the registered client and stored refresh token", form) + } + + // The refresh must not drop what Login stored, or the next one would fail. + saved, ok, err := store.LoadForServer(server) + if err != nil || !ok { + t.Fatalf("LoadForServer() = %v, %v", ok, err) + } + if saved.RefreshToken != "refresh-rotated" || saved.TokenEndpoint != fixture.URL+"/token" || !saved.ProtectedTokenEndpoint || + saved.ClientID != "registered-client" || saved.ClientSecret != "registered-secret" { + t.Fatalf("refreshed token lost the stored endpoint or client: %#v", saved) + } + if _, err := source.Refresh(context.Background()); err != nil { + t.Fatalf("second Refresh() error = %v", err) + } +} + +func TestRefreshSettingsConfigOverridesStored(t *testing.T) { + cfg, client, err := refreshSettings(context.Background(), http.DefaultClient, "https://mcp.example.com/mcp", + OAuthConfig{ClientID: "configured-client", ClientSecret: "configured-secret", TokenEndpoint: "https://configured.example.com/token"}, + StoredToken{TokenEndpoint: "https://stored.example.com/token", ProtectedTokenEndpoint: true, ClientID: "stored-client", ClientSecret: "stored-secret"}) + if err != nil { + t.Fatalf("refreshSettings() error = %v", err) + } + if cfg.TokenEndpoint != "https://configured.example.com/token" || cfg.ClientID != "configured-client" || cfg.ClientSecret != "configured-secret" { + t.Fatalf("config did not win over stored values: %#v", cfg) + } + if client != http.DefaultClient { + t.Fatal("an explicitly configured endpoint must not switch to the protected-discovery client") + } +} + +func TestRefreshSettingsKeepsConfiguredSecretWithStoredClientID(t *testing.T) { + cfg, _, err := refreshSettings(context.Background(), http.DefaultClient, "https://mcp.example.com/mcp", + OAuthConfig{ClientSecret: "configured-secret"}, + StoredToken{ClientID: "stored-client", ClientSecret: "stored-secret"}) + if err != nil { + t.Fatalf("refreshSettings() error = %v", err) + } + if cfg.ClientID != "stored-client" || cfg.ClientSecret != "configured-secret" { + t.Fatalf("client = %q/%q, want stored id with the configured secret", cfg.ClientID, cfg.ClientSecret) + } +} + +func TestRefreshSettingsRejectsUnsafeStoredEndpoint(t *testing.T) { + for _, tc := range []struct { + name string + stored StoredToken + want error + }{ + { + name: "cleartext endpoint", + stored: StoredToken{TokenEndpoint: "http://auth.example.com/token", ClientID: "c"}, + want: oauth.ErrInsecureTokenEndpoint, + }, + { + name: "advertised loopback endpoint", + stored: StoredToken{TokenEndpoint: "https://127.0.0.1/token", ProtectedTokenEndpoint: true, ClientID: "c"}, + want: errUnsafeOAuthDiscoveryTarget, + }, + { + name: "advertised private endpoint", + stored: StoredToken{TokenEndpoint: "https://10.0.0.5/token", ProtectedTokenEndpoint: true, ClientID: "c"}, + want: errUnsafeOAuthDiscoveryTarget, + }, + } { + t.Run(tc.name, func(t *testing.T) { + _, _, err := refreshSettings(context.Background(), http.DefaultClient, "https://mcp.example.com/mcp", OAuthConfig{}, tc.stored) + if !errors.Is(err, tc.want) { + t.Fatalf("refreshSettings() error = %v, want %v", err, tc.want) + } + }) + } +} + +func TestStoreTokenSourceRefreshDoesNotPostToUnsafeStoredEndpoint(t *testing.T) { + var hits atomic.Int64 + loopback := httptest.NewServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { hits.Add(1) })) + defer loopback.Close() + + // A public MCP resource must not be able to point refresh at loopback. + server := Server{Name: "public", URL: "https://mcp.example.com/mcp", Identity: testServerIdentity, Auth: ServerAuthOAuth} + store, err := NewTokenStore(TokenStoreOptions{FilePath: filepath.Join(t.TempDir(), "oauth-tokens.json")}) + if err != nil { + t.Fatalf("NewTokenStore() error = %v", err) + } + if err := store.SaveForServer(server, StoredToken{ + AccessToken: "a", RefreshToken: "r", ClientID: "c", + TokenEndpoint: loopback.URL + "/token", ProtectedTokenEndpoint: true, + }); err != nil { + t.Fatalf("SaveForServer() error = %v", err) + } + source := &storeTokenSource{server: server, store: store, httpClient: http.DefaultClient, now: time.Now} + if _, err := source.Refresh(context.Background()); !errors.Is(err, errUnsafeOAuthDiscoveryTarget) { + t.Fatalf("Refresh() error = %v, want the discovery policy refusal", err) + } + if hits.Load() != 0 { + t.Fatalf("loopback endpoint received %d request(s)", hits.Load()) + } +} + +func TestStoredTokenPersistsDiscoveryFields(t *testing.T) { + store, err := NewTokenStore(TokenStoreOptions{FilePath: filepath.Join(t.TempDir(), "oauth-tokens.json")}) + if err != nil { + t.Fatalf("NewTokenStore() error = %v", err) + } + want := StoredToken{ + AccessToken: "a", RefreshToken: "r", + TokenEndpoint: "https://auth.example.com/token", ProtectedTokenEndpoint: true, + ClientID: "client", ClientSecret: "secret", + } + if err := store.Save("demo", want); err != nil { + t.Fatalf("Save() error = %v", err) + } + got, ok, err := store.Load("demo") + if err != nil || !ok { + t.Fatalf("Load() = %v, %v", ok, err) + } + if got.TokenEndpoint != want.TokenEndpoint || !got.ProtectedTokenEndpoint || got.ClientID != want.ClientID || got.ClientSecret != want.ClientSecret { + t.Fatalf("round trip lost discovery fields: %#v", got) + } + statuses, err := store.Status() + if err != nil { + t.Fatalf("Status() error = %v", err) + } + encoded, _ := json.Marshal(statuses) + for _, secret := range []string{"secret", "https://auth.example.com/token"} { + if strings.Contains(string(encoded), secret) { + t.Fatalf("status output leaked %q: %s", secret, encoded) + } + } +} diff --git a/internal/mcp/oauth_store.go b/internal/mcp/oauth_store.go index ce730182a..18aea0a2e 100644 --- a/internal/mcp/oauth_store.go +++ b/internal/mcp/oauth_store.go @@ -29,6 +29,15 @@ type StoredToken struct { TokenType string `json:"token_type,omitempty"` Scopes []string `json:"scopes,omitempty"` ExpiresAt time.Time `json:"expires_at,omitempty"` + // TokenEndpoint, ClientID and ClientSecret hold what Login discovered or + // dynamically registered, so Refresh works for the standard `auth: oauth` + // setup whose config carries none of them. ClientSecret is sensitive. + // ProtectedTokenEndpoint is set when the endpoint came from server-advertised + // discovery and must pass the public-network policy again before use. + TokenEndpoint string `json:"token_endpoint,omitempty"` + ProtectedTokenEndpoint bool `json:"protected_token_endpoint,omitempty"` + ClientID string `json:"client_id,omitempty"` + ClientSecret string `json:"client_secret,omitempty"` } // TokenStatus is a redaction-safe summary of a stored token. It deliberately @@ -393,6 +402,11 @@ func storedToOAuth(s StoredToken) oauth.Token { TokenType: s.TokenType, Scopes: s.Scopes, ExpiresAt: s.ExpiresAt, + + TokenEndpoint: s.TokenEndpoint, + ProtectedTokenEndpoint: s.ProtectedTokenEndpoint, + ClientID: s.ClientID, + ClientSecret: s.ClientSecret, } } diff --git a/internal/oauth/oauth.go b/internal/oauth/oauth.go index 4938e7bf3..1d2f59111 100644 --- a/internal/oauth/oauth.go +++ b/internal/oauth/oauth.go @@ -45,6 +45,16 @@ type Token struct { // `chatgpt_account_id`) that the request path needs as headers. Treated as // sensitive like the access token: never logged, persisted 0600. IDToken string `json:"id_token,omitempty"` + // TokenEndpoint, ClientID and ClientSecret record what a login learned that + // config does not carry (a discovered token endpoint, a dynamically + // registered client) so a later refresh can reuse it. ClientSecret is + // sensitive like the tokens. ProtectedTokenEndpoint marks an endpoint + // advertised by the server, which must be re-checked against the public + // network policy before a refresh posts credentials to it. + TokenEndpoint string `json:"token_endpoint,omitempty"` + ProtectedTokenEndpoint bool `json:"protected_token_endpoint,omitempty"` + ClientID string `json:"client_id,omitempty"` + ClientSecret string `json:"client_secret,omitempty"` } // Expired reports whether the token has an expiry that is at or before now.