From e48909215430b39b91ede707cf5240b27a21bb41 Mon Sep 17 00:00:00 2001 From: Michael Suchacz <203725896+ibetitsmike@users.noreply.github.com> Date: Thu, 16 Jul 2026 13:43:05 +0200 Subject: [PATCH] feat: handle revoked OAuth grants for MCP servers gracefully (#27264) Closes [CODAGT-792](https://linear.app/codercom/issue/CODAGT-792/handle-revoked-oauth-grants-for-mcp-servers-gracefully). When a user revokes an upstream OAuth grant for an MCP server used by Coder Agents, Coder kept treating the cached token as valid: `invalid_grant` refresh failures were logged and swallowed, the dead bearer token kept being attached, the list endpoints re-attempted the refresh on every call, and the UI kept showing the server as authenticated. ## Changes Backend, mirroring the `external_auth_links` prior art: - New migration adds `mcp_server_user_tokens.oauth_refresh_failure_reason`. `UpsertMCPServerUserToken` clears it, so completing the OAuth flow again recovers the row. - New `MarkMCPServerUserTokenRefreshFailure` query records the failure and clears all token material, guarded by an `updated_at` optimistic lock so a stale failure never clobbers a concurrently refreshed token (on a lock miss the winner's row is used). - `mcpclient.IsPermanentRefreshError` classifies `*oauth2.RetrieveError` codes: only `invalid_grant` and `bad_refresh_token` are permanent. Client/config errors (`invalid_client`, `unauthorized_client`, ...) stay transient for the user row since reconnecting cannot fix them. - chatd token refresh and the MCP list/get endpoints persist permanent failures, return cleared tokens for the in-flight request, and skip provider calls for already-failed rows. - `buildAuthHeaders` no longer attaches an Authorization header for failed tokens, so chat degrades by omitting that server's tools instead of sending a dead bearer. API and UI: - No new API surface. A permanently failed token simply reports `auth_connected: false`, so the existing "Auth" button and "Not authenticated" tooltip appear and the user re-runs the same OAuth flow to recover. An earlier revision added an `auth_status` enum (`connected` / `not_connected` / `reconnect_required`) with a dedicated "Reconnect" button; it was collapsed to keep the API minimal since both states lead to the identical re-auth action. Out of scope (follow-up): typed 401-on-connect detection and forced refresh. mcp-go exposes no stable typed 401 signal in the static-header path, so a revocation while the access token still looks valid locally stays undetected until expiry triggers a refresh. ## Testing - Unit and integration tests: classifier, chatd refresh paths (permanent/transient/race/persist-failure), API endpoints (revoked, transient, no-retry caching, re-auth recovery, stale-lock), dbauthz, dbcrypt, migrations. - Dogfood UAT against a dev instance with a mock IdP returning `invalid_grant`: revoked grant detected on refresh and persisted once (no repeated IdP calls), chat with the revoked server selected completes with the server's tools omitted, and re-auth restores the connected state. > This PR was authored by Mux, working on Mike's behalf. --- coderd/database/dbauthz/dbauthz.go | 7 + coderd/database/dbauthz/dbauthz_test.go | 10 + coderd/database/dbmetrics/querymetrics.go | 8 + coderd/database/dbmock/dbmock.go | 15 ++ coderd/database/dump.sql | 3 +- .../000544_mcp_token_refresh_failure.down.sql | 3 + .../000544_mcp_token_refresh_failure.up.sql | 3 + coderd/database/models.go | 23 +- coderd/database/querier.go | 6 + coderd/database/queries.sql.go | 60 ++++- coderd/database/queries/mcpserverconfigs.sql | 24 ++ coderd/mcp.go | 80 ++++++- coderd/mcp_test.go | 190 ++++++++++++++++ coderd/x/chatd/chatd.go | 72 ++++++ coderd/x/chatd/mcp_refresh_internal_test.go | 215 ++++++++++++++++++ coderd/x/chatd/mcpclient/export_test.go | 4 + coderd/x/chatd/mcpclient/mcpclient.go | 47 ++++ coderd/x/chatd/mcpclient/refresh_test.go | 136 +++++++++++ enterprise/dbcrypt/dbcrypt.go | 14 ++ enterprise/dbcrypt/dbcrypt_internal_test.go | 18 ++ 20 files changed, 913 insertions(+), 25 deletions(-) create mode 100644 coderd/database/migrations/000544_mcp_token_refresh_failure.down.sql create mode 100644 coderd/database/migrations/000544_mcp_token_refresh_failure.up.sql create mode 100644 coderd/x/chatd/mcp_refresh_internal_test.go create mode 100644 coderd/x/chatd/mcpclient/refresh_test.go diff --git a/coderd/database/dbauthz/dbauthz.go b/coderd/database/dbauthz/dbauthz.go index fd9e9d685b..58094fdc65 100644 --- a/coderd/database/dbauthz/dbauthz.go +++ b/coderd/database/dbauthz/dbauthz.go @@ -6869,6 +6869,13 @@ func (q *querier) MarkChatsContextDirtyByAgent(ctx context.Context, arg database return q.db.MarkChatsContextDirtyByAgent(ctx, arg) } +func (q *querier) MarkMCPServerUserTokenRefreshFailure(ctx context.Context, arg database.MarkMCPServerUserTokenRefreshFailureParams) (database.MCPServerUserToken, error) { + if err := q.authorizeContext(ctx, policy.ActionUpdate, rbac.ResourceDeploymentConfig); err != nil { + return database.MCPServerUserToken{}, err + } + return q.db.MarkMCPServerUserTokenRefreshFailure(ctx, arg) +} + func (q *querier) OIDCClaimFieldValues(ctx context.Context, args database.OIDCClaimFieldValuesParams) ([]string, error) { resource := rbac.ResourceIdpsyncSettings if args.OrganizationID != uuid.Nil { diff --git a/coderd/database/dbauthz/dbauthz_test.go b/coderd/database/dbauthz/dbauthz_test.go index 0b1c2b953b..fc05ce9411 100644 --- a/coderd/database/dbauthz/dbauthz_test.go +++ b/coderd/database/dbauthz/dbauthz_test.go @@ -1925,6 +1925,16 @@ func (s *MethodTestSuite) TestChats() { dbm.EXPECT().UpsertMCPServerUserToken(gomock.Any(), arg).Return(token, nil).AnyTimes() check.Args(arg).Asserts(rbac.ResourceDeploymentConfig, policy.ActionUpdate).Returns(token) })) + s.Run("MarkMCPServerUserTokenRefreshFailure", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { + token := testutil.Fake(s.T(), faker, database.MCPServerUserToken{}) + arg := database.MarkMCPServerUserTokenRefreshFailureParams{ + ID: token.ID, + UpdatedAt: token.UpdatedAt, + OauthRefreshFailureReason: "invalid_grant", + } + dbm.EXPECT().MarkMCPServerUserTokenRefreshFailure(gomock.Any(), arg).Return(token, nil).AnyTimes() + check.Args(arg).Asserts(rbac.ResourceDeploymentConfig, policy.ActionUpdate).Returns(token) + })) } func (s *MethodTestSuite) TestFile() { diff --git a/coderd/database/dbmetrics/querymetrics.go b/coderd/database/dbmetrics/querymetrics.go index 145d7ff582..6a964258b0 100644 --- a/coderd/database/dbmetrics/querymetrics.go +++ b/coderd/database/dbmetrics/querymetrics.go @@ -4841,6 +4841,14 @@ func (m queryMetricsStore) MarkChatsContextDirtyByAgent(ctx context.Context, arg return r0, r1 } +func (m queryMetricsStore) MarkMCPServerUserTokenRefreshFailure(ctx context.Context, arg database.MarkMCPServerUserTokenRefreshFailureParams) (database.MCPServerUserToken, error) { + start := time.Now() + r0, r1 := m.s.MarkMCPServerUserTokenRefreshFailure(ctx, arg) + m.queryLatencies.WithLabelValues("MarkMCPServerUserTokenRefreshFailure").Observe(time.Since(start).Seconds()) + m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "MarkMCPServerUserTokenRefreshFailure").Inc() + return r0, r1 +} + func (m queryMetricsStore) OIDCClaimFieldValues(ctx context.Context, arg database.OIDCClaimFieldValuesParams) ([]string, error) { start := time.Now() r0, r1 := m.s.OIDCClaimFieldValues(ctx, arg) diff --git a/coderd/database/dbmock/dbmock.go b/coderd/database/dbmock/dbmock.go index b2b7c28e02..f9f5738db6 100644 --- a/coderd/database/dbmock/dbmock.go +++ b/coderd/database/dbmock/dbmock.go @@ -9114,6 +9114,21 @@ func (mr *MockStoreMockRecorder) MarkChatsContextDirtyByAgent(ctx, arg any) *gom return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "MarkChatsContextDirtyByAgent", reflect.TypeOf((*MockStore)(nil).MarkChatsContextDirtyByAgent), ctx, arg) } +// MarkMCPServerUserTokenRefreshFailure mocks base method. +func (m *MockStore) MarkMCPServerUserTokenRefreshFailure(ctx context.Context, arg database.MarkMCPServerUserTokenRefreshFailureParams) (database.MCPServerUserToken, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "MarkMCPServerUserTokenRefreshFailure", ctx, arg) + ret0, _ := ret[0].(database.MCPServerUserToken) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// MarkMCPServerUserTokenRefreshFailure indicates an expected call of MarkMCPServerUserTokenRefreshFailure. +func (mr *MockStoreMockRecorder) MarkMCPServerUserTokenRefreshFailure(ctx, arg any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "MarkMCPServerUserTokenRefreshFailure", reflect.TypeOf((*MockStore)(nil).MarkMCPServerUserTokenRefreshFailure), ctx, arg) +} + // OIDCClaimFieldValues mocks base method. func (m *MockStore) OIDCClaimFieldValues(ctx context.Context, arg database.OIDCClaimFieldValuesParams) ([]string, error) { m.ctrl.T.Helper() diff --git a/coderd/database/dump.sql b/coderd/database/dump.sql index 8c9855b2ed..081db111f3 100644 --- a/coderd/database/dump.sql +++ b/coderd/database/dump.sql @@ -2490,7 +2490,8 @@ CREATE TABLE mcp_server_user_tokens ( token_type text DEFAULT 'Bearer'::text NOT NULL, expiry timestamp with time zone, created_at timestamp with time zone DEFAULT now() NOT NULL, - updated_at timestamp with time zone DEFAULT now() NOT NULL + updated_at timestamp with time zone DEFAULT now() NOT NULL, + oauth_refresh_failure_reason text DEFAULT ''::text NOT NULL ); CREATE TABLE notification_messages ( diff --git a/coderd/database/migrations/000544_mcp_token_refresh_failure.down.sql b/coderd/database/migrations/000544_mcp_token_refresh_failure.down.sql new file mode 100644 index 0000000000..86db3baa42 --- /dev/null +++ b/coderd/database/migrations/000544_mcp_token_refresh_failure.down.sql @@ -0,0 +1,3 @@ +ALTER TABLE mcp_server_user_tokens + DROP COLUMN oauth_refresh_failure_reason +; diff --git a/coderd/database/migrations/000544_mcp_token_refresh_failure.up.sql b/coderd/database/migrations/000544_mcp_token_refresh_failure.up.sql new file mode 100644 index 0000000000..d300b03864 --- /dev/null +++ b/coderd/database/migrations/000544_mcp_token_refresh_failure.up.sql @@ -0,0 +1,3 @@ +ALTER TABLE mcp_server_user_tokens + ADD COLUMN oauth_refresh_failure_reason TEXT NOT NULL DEFAULT '' +; diff --git a/coderd/database/models.go b/coderd/database/models.go index 7edc3e6f4c..263d1f0fbc 100644 --- a/coderd/database/models.go +++ b/coderd/database/models.go @@ -5438,17 +5438,18 @@ type MCPServerConfig struct { } type MCPServerUserToken struct { - ID uuid.UUID `db:"id" json:"id"` - MCPServerConfigID uuid.UUID `db:"mcp_server_config_id" json:"mcp_server_config_id"` - UserID uuid.UUID `db:"user_id" json:"user_id"` - AccessToken string `db:"access_token" json:"access_token"` - AccessTokenKeyID sql.NullString `db:"access_token_key_id" json:"access_token_key_id"` - RefreshToken string `db:"refresh_token" json:"refresh_token"` - RefreshTokenKeyID sql.NullString `db:"refresh_token_key_id" json:"refresh_token_key_id"` - TokenType string `db:"token_type" json:"token_type"` - Expiry sql.NullTime `db:"expiry" json:"expiry"` - CreatedAt time.Time `db:"created_at" json:"created_at"` - UpdatedAt time.Time `db:"updated_at" json:"updated_at"` + ID uuid.UUID `db:"id" json:"id"` + MCPServerConfigID uuid.UUID `db:"mcp_server_config_id" json:"mcp_server_config_id"` + UserID uuid.UUID `db:"user_id" json:"user_id"` + AccessToken string `db:"access_token" json:"access_token"` + AccessTokenKeyID sql.NullString `db:"access_token_key_id" json:"access_token_key_id"` + RefreshToken string `db:"refresh_token" json:"refresh_token"` + RefreshTokenKeyID sql.NullString `db:"refresh_token_key_id" json:"refresh_token_key_id"` + TokenType string `db:"token_type" json:"token_type"` + Expiry sql.NullTime `db:"expiry" json:"expiry"` + CreatedAt time.Time `db:"created_at" json:"created_at"` + UpdatedAt time.Time `db:"updated_at" json:"updated_at"` + OauthRefreshFailureReason string `db:"oauth_refresh_failure_reason" json:"oauth_refresh_failure_reason"` } type NotificationMessage struct { diff --git a/coderd/database/querier.go b/coderd/database/querier.go index dbc70a813e..cd0625c0ff 100644 --- a/coderd/database/querier.go +++ b/coderd/database/querier.go @@ -1205,6 +1205,12 @@ type sqlcQuerier interface { // re-pins it. Returns the chats that transitioned so the caller can // emit watch events after the transaction commits. MarkChatsContextDirtyByAgent(ctx context.Context, arg MarkChatsContextDirtyByAgentParams) ([]MarkChatsContextDirtyByAgentRow, error) + // Records a permanent refresh failure (e.g. revoked grant) and clears + // the dead token material so it is never attached to a request again. + // The updated_at predicate provides optimistic concurrency: if another + // request refreshed or replaced the token since it was read, this + // update matches zero rows and returns sql.ErrNoRows. + MarkMCPServerUserTokenRefreshFailure(ctx context.Context, arg MarkMCPServerUserTokenRefreshFailureParams) (MCPServerUserToken, error) OIDCClaimFieldValues(ctx context.Context, arg OIDCClaimFieldValuesParams) ([]string, error) // OIDCClaimFields returns a list of distinct keys in the the merged_claims fields. // This query is used to generate the list of available sync fields for idp sync settings. diff --git a/coderd/database/queries.sql.go b/coderd/database/queries.sql.go index 574c803aa5..6271762ca7 100644 --- a/coderd/database/queries.sql.go +++ b/coderd/database/queries.sql.go @@ -16864,7 +16864,7 @@ func (q *sqlQuerier) GetMCPServerConfigsByIDs(ctx context.Context, ids []uuid.UU const getMCPServerUserToken = `-- name: GetMCPServerUserToken :one SELECT - id, mcp_server_config_id, user_id, access_token, access_token_key_id, refresh_token, refresh_token_key_id, token_type, expiry, created_at, updated_at + id, mcp_server_config_id, user_id, access_token, access_token_key_id, refresh_token, refresh_token_key_id, token_type, expiry, created_at, updated_at, oauth_refresh_failure_reason FROM mcp_server_user_tokens WHERE @@ -16892,13 +16892,14 @@ func (q *sqlQuerier) GetMCPServerUserToken(ctx context.Context, arg GetMCPServer &i.Expiry, &i.CreatedAt, &i.UpdatedAt, + &i.OauthRefreshFailureReason, ) return i, err } const getMCPServerUserTokensByUserID = `-- name: GetMCPServerUserTokensByUserID :many SELECT - id, mcp_server_config_id, user_id, access_token, access_token_key_id, refresh_token, refresh_token_key_id, token_type, expiry, created_at, updated_at + id, mcp_server_config_id, user_id, access_token, access_token_key_id, refresh_token, refresh_token_key_id, token_type, expiry, created_at, updated_at, oauth_refresh_failure_reason FROM mcp_server_user_tokens WHERE @@ -16926,6 +16927,7 @@ func (q *sqlQuerier) GetMCPServerUserTokensByUserID(ctx context.Context, userID &i.Expiry, &i.CreatedAt, &i.UpdatedAt, + &i.OauthRefreshFailureReason, ); err != nil { return nil, err } @@ -17098,6 +17100,54 @@ func (q *sqlQuerier) InsertMCPServerConfig(ctx context.Context, arg InsertMCPSer return i, err } +const markMCPServerUserTokenRefreshFailure = `-- name: MarkMCPServerUserTokenRefreshFailure :one +UPDATE mcp_server_user_tokens +SET + access_token = '', + access_token_key_id = NULL, + refresh_token = '', + refresh_token_key_id = NULL, + expiry = NULL, + oauth_refresh_failure_reason = $1::text, + updated_at = NOW() +WHERE + id = $2::uuid + AND updated_at = $3::timestamptz +RETURNING + id, mcp_server_config_id, user_id, access_token, access_token_key_id, refresh_token, refresh_token_key_id, token_type, expiry, created_at, updated_at, oauth_refresh_failure_reason +` + +type MarkMCPServerUserTokenRefreshFailureParams struct { + OauthRefreshFailureReason string `db:"oauth_refresh_failure_reason" json:"oauth_refresh_failure_reason"` + ID uuid.UUID `db:"id" json:"id"` + UpdatedAt time.Time `db:"updated_at" json:"updated_at"` +} + +// Records a permanent refresh failure (e.g. revoked grant) and clears +// the dead token material so it is never attached to a request again. +// The updated_at predicate provides optimistic concurrency: if another +// request refreshed or replaced the token since it was read, this +// update matches zero rows and returns sql.ErrNoRows. +func (q *sqlQuerier) MarkMCPServerUserTokenRefreshFailure(ctx context.Context, arg MarkMCPServerUserTokenRefreshFailureParams) (MCPServerUserToken, error) { + row := q.db.QueryRowContext(ctx, markMCPServerUserTokenRefreshFailure, arg.OauthRefreshFailureReason, arg.ID, arg.UpdatedAt) + var i MCPServerUserToken + err := row.Scan( + &i.ID, + &i.MCPServerConfigID, + &i.UserID, + &i.AccessToken, + &i.AccessTokenKeyID, + &i.RefreshToken, + &i.RefreshTokenKeyID, + &i.TokenType, + &i.Expiry, + &i.CreatedAt, + &i.UpdatedAt, + &i.OauthRefreshFailureReason, + ) + return i, err +} + const updateMCPServerConfig = `-- name: UpdateMCPServerConfig :one UPDATE mcp_server_configs @@ -17258,9 +17308,12 @@ ON CONFLICT (mcp_server_config_id, user_id) DO UPDATE SET refresh_token_key_id = $6::text, token_type = $7::text, expiry = $8::timestamptz, + -- New token material means the user re-authenticated, so any + -- cached permanent refresh failure no longer applies. + oauth_refresh_failure_reason = '', updated_at = NOW() RETURNING - id, mcp_server_config_id, user_id, access_token, access_token_key_id, refresh_token, refresh_token_key_id, token_type, expiry, created_at, updated_at + id, mcp_server_config_id, user_id, access_token, access_token_key_id, refresh_token, refresh_token_key_id, token_type, expiry, created_at, updated_at, oauth_refresh_failure_reason ` type UpsertMCPServerUserTokenParams struct { @@ -17298,6 +17351,7 @@ func (q *sqlQuerier) UpsertMCPServerUserToken(ctx context.Context, arg UpsertMCP &i.Expiry, &i.CreatedAt, &i.UpdatedAt, + &i.OauthRefreshFailureReason, ) return i, err } diff --git a/coderd/database/queries/mcpserverconfigs.sql b/coderd/database/queries/mcpserverconfigs.sql index 3d05a2b102..be7c3f6622 100644 --- a/coderd/database/queries/mcpserverconfigs.sql +++ b/coderd/database/queries/mcpserverconfigs.sql @@ -200,10 +200,34 @@ ON CONFLICT (mcp_server_config_id, user_id) DO UPDATE SET refresh_token_key_id = sqlc.narg('refresh_token_key_id')::text, token_type = @token_type::text, expiry = sqlc.narg('expiry')::timestamptz, + -- New token material means the user re-authenticated, so any + -- cached permanent refresh failure no longer applies. + oauth_refresh_failure_reason = '', updated_at = NOW() RETURNING *; +-- name: MarkMCPServerUserTokenRefreshFailure :one +-- Records a permanent refresh failure (e.g. revoked grant) and clears +-- the dead token material so it is never attached to a request again. +-- The updated_at predicate provides optimistic concurrency: if another +-- request refreshed or replaced the token since it was read, this +-- update matches zero rows and returns sql.ErrNoRows. +UPDATE mcp_server_user_tokens +SET + access_token = '', + access_token_key_id = NULL, + refresh_token = '', + refresh_token_key_id = NULL, + expiry = NULL, + oauth_refresh_failure_reason = @oauth_refresh_failure_reason::text, + updated_at = NOW() +WHERE + id = @id::uuid + AND updated_at = @updated_at::timestamptz +RETURNING + *; + -- name: DeleteMCPServerUserToken :exec DELETE FROM mcp_server_user_tokens diff --git a/coderd/mcp.go b/coderd/mcp.go index 3e0a5829f7..8cea933369 100644 --- a/coderd/mcp.go +++ b/coderd/mcp.go @@ -208,8 +208,6 @@ func (api *API) listMCPServerConfigs(rw http.ResponseWriter, r *http.Request) { } if config.AuthType == "oauth2" { sdkConfig.AuthConnected = tokenMap[config.ID] - } else { - sdkConfig.AuthConnected = true } resp = append(resp, sdkConfig) } @@ -537,8 +535,6 @@ func (api *API) getMCPServerConfig(rw http.ResponseWriter, r *http.Request) { break } } - } else { - sdkConfig.AuthConnected = true } httpapi.Write(ctx, rw, http.StatusOK, sdkConfig) @@ -1167,12 +1163,12 @@ func (api *API) mcpServerOAuth2Disconnect(rw http.ResponseWriter, r *http.Reques rw.WriteHeader(http.StatusNoContent) } -// parseMCPServerConfigID extracts the MCP server config UUID from the -// "mcpServer" path parameter. // refreshMCPUserToken attempts to refresh an expired OAuth2 token // for the given MCP server config. Returns true when the token is // valid (either still fresh or successfully refreshed), false when -// the token is expired and cannot be refreshed. +// the token is expired and cannot be refreshed. Permanent refresh +// failures (e.g. revoked grants) are persisted so subsequent calls +// skip the provider without a network call. func (api *API) refreshMCPUserToken( ctx context.Context, cfg database.MCPServerConfig, @@ -1182,9 +1178,12 @@ func (api *API) refreshMCPUserToken( if cfg.AuthType != "oauth2" { return true } + if tok.OauthRefreshFailureReason != "" { + return false + } if tok.RefreshToken == "" { - // No refresh token — consider connected only if not - // expired (or no expiry set). + // No refresh token; connected only if not expired (or no + // expiry set). return !tok.Expiry.Valid || tok.Expiry.Time.After(time.Now()) } @@ -1194,7 +1193,11 @@ func (api *API) refreshMCPUserToken( slog.F("server_slug", cfg.Slug), slog.Error(err), ) - // Refresh failed — token is dead. + if mcpclient.IsPermanentRefreshError(err) { + return api.markMCPTokenRefreshFailure(ctx, cfg, tok, err) + } + // Transient failure; the token is unusable right now but a + // later refresh may succeed. return false } @@ -1230,6 +1233,59 @@ func (api *API) refreshMCPUserToken( return true } +// markMCPTokenRefreshFailure persists a permanent refresh failure so +// later status checks skip the provider. The updated_at optimistic +// lock loses to concurrent refreshes: in that case the winner's row +// determines whether the token is still usable. +func (api *API) markMCPTokenRefreshFailure( + ctx context.Context, + cfg database.MCPServerConfig, + tok database.MCPServerUserToken, + refreshErr error, +) bool { + //nolint:gocritic // Need system-level write access to persist + // the refresh failure. + _, err := api.Database.MarkMCPServerUserTokenRefreshFailure( + dbauthz.AsSystemRestricted(ctx), + database.MarkMCPServerUserTokenRefreshFailureParams{ + ID: tok.ID, + UpdatedAt: tok.UpdatedAt, + OauthRefreshFailureReason: mcpclient.RefreshFailureReason(refreshErr), + }, + ) + if err == nil { + return false + } + + if xerrors.Is(err, sql.ErrNoRows) { + // A concurrent request updated the token after we read it; + // report its state instead of poisoning the fresh token. + //nolint:gocritic // Need system-level read access to load + // the concurrently updated token. + current, readErr := api.Database.GetMCPServerUserToken( + dbauthz.AsSystemRestricted(ctx), + database.GetMCPServerUserTokenParams{ + MCPServerConfigID: tok.MCPServerConfigID, + UserID: tok.UserID, + }, + ) + if readErr == nil { + return current.OauthRefreshFailureReason == "" && + current.AccessToken != "" && + (!current.Expiry.Valid || current.Expiry.Time.After(time.Now())) + } + err = readErr + } + + api.Logger.Warn(ctx, "failed to persist MCP oauth2 refresh failure", + slog.F("server_slug", cfg.Slug), + slog.Error(err), + ) + return false +} + +// parseMCPServerConfigID extracts the MCP server config UUID from the +// "mcpServer" path parameter. func parseMCPServerConfigID(rw http.ResponseWriter, r *http.Request) (uuid.UUID, bool) { mcpServerID, err := uuid.Parse(chi.URLParam(r, "mcpServer")) if err != nil { @@ -1279,6 +1335,10 @@ func convertMCPServerConfig(config database.MCPServerConfig) codersdk.MCPServerC ForwardCoderHeaders: config.ForwardCoderHeaders, CreatedAt: config.CreatedAt, UpdatedAt: config.UpdatedAt, + + // Default per-user auth state. Handlers that know the + // calling user's token state (list/get) overwrite this. + AuthConnected: config.AuthType != "oauth2", } } diff --git a/coderd/mcp_test.go b/coderd/mcp_test.go index dde85f12e7..c3e7721207 100644 --- a/coderd/mcp_test.go +++ b/coderd/mcp_test.go @@ -2,18 +2,23 @@ package coderd_test import ( "crypto/sha256" + "database/sql" "encoding/base64" "encoding/json" "net/http" "net/http/httptest" "strings" + "sync/atomic" "testing" + "time" "github.com/google/uuid" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "github.com/coder/coder/v2/coderd/coderdtest" + "github.com/coder/coder/v2/coderd/database" + "github.com/coder/coder/v2/coderd/database/dbauthz" "github.com/coder/coder/v2/codersdk" "github.com/coder/coder/v2/testutil" ) @@ -1944,3 +1949,188 @@ func TestMCPOAuth2DiscoveryEdgeCases(t *testing.T) { require.True(t, created.HasOAuth2Secret) }) } + +func TestMCPServerConfigsRevokedGrant(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitLong) + providerKeys := coderdtest.FakeOpenAICompatProviderAPIKeys(t) + adminClient, db := coderdtest.NewWithDatabase(t, &coderdtest.Options{ + DeploymentValues: mcpDeploymentValues(t), + ChatProviderAPIKeys: &providerKeys, + }) + firstUser := coderdtest.CreateFirstUser(t, adminClient) + memberClient, member := coderdtest.CreateAnotherUser(t, adminClient, firstUser.OrganizationID) + + var tokenEndpointHits atomic.Int64 + tokenSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + tokenEndpointHits.Add(1) + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusBadRequest) + _, _ = w.Write([]byte(`{"error":"invalid_grant","error_description":"grant revoked"}`)) + })) + t.Cleanup(tokenSrv.Close) + + created, err := adminClient.CreateMCPServerConfig(ctx, codersdk.CreateMCPServerConfigRequest{ + DisplayName: "Revoked Server", + Slug: "revoked-server", + Transport: "streamable_http", + URL: "https://mcp.example.com/v1", + AuthType: "oauth2", + OAuth2ClientID: "cid", + OAuth2AuthURL: "https://auth.example.com/authorize", + OAuth2TokenURL: tokenSrv.URL, + Availability: "default_on", + Enabled: true, + ToolAllowList: []string{}, + ToolDenyList: []string{}, + }) + require.NoError(t, err) + require.False(t, created.AuthConnected) + + // Seed an expired token whose refresh the provider rejects with + // invalid_grant. + //nolint:gocritic // Seeding test state requires system access. + seeded, err := db.UpsertMCPServerUserToken(dbauthz.AsSystemRestricted(ctx), database.UpsertMCPServerUserTokenParams{ + MCPServerConfigID: created.ID, + UserID: member.ID, + AccessToken: "expired-access", + RefreshToken: "dead-refresh", + TokenType: "Bearer", + Expiry: sql.NullTime{Time: time.Now().Add(-time.Hour), Valid: true}, + }) + require.NoError(t, err) + + // First list: the refresh fails permanently, so the server is + // reported as not connected and the failure is persisted. + configs, err := memberClient.MCPServerConfigs(ctx) + require.NoError(t, err) + require.Len(t, configs, 1) + require.False(t, configs[0].AuthConnected) + // The oauth2 package may probe both client auth styles, so the + // exact count varies; what matters is that it never grows again. + hitsAfterFirstList := tokenEndpointHits.Load() + require.Positive(t, hitsAfterFirstList) + + //nolint:gocritic // Verifying persisted state requires system access. + row, err := db.GetMCPServerUserToken(dbauthz.AsSystemRestricted(ctx), database.GetMCPServerUserTokenParams{ + MCPServerConfigID: created.ID, + UserID: member.ID, + }) + require.NoError(t, err) + require.Empty(t, row.AccessToken) + require.Empty(t, row.RefreshToken) + require.False(t, row.Expiry.Valid) + require.Contains(t, row.OauthRefreshFailureReason, "invalid_grant") + + // Second list: the cached failure short-circuits, so the provider + // is not called again. + configs, err = memberClient.MCPServerConfigs(ctx) + require.NoError(t, err) + require.Len(t, configs, 1) + require.False(t, configs[0].AuthConnected) + require.Equal(t, hitsAfterFirstList, tokenEndpointHits.Load()) + + // The single-config endpoint agrees. + single, err := memberClient.MCPServerConfigByID(ctx, created.ID) + require.NoError(t, err) + require.False(t, single.AuthConnected) + require.Equal(t, hitsAfterFirstList, tokenEndpointHits.Load()) + + // A stale optimistic-lock update must not clobber the row. + //nolint:gocritic // Exercising the query requires system access. + _, err = db.MarkMCPServerUserTokenRefreshFailure(dbauthz.AsSystemRestricted(ctx), database.MarkMCPServerUserTokenRefreshFailureParams{ + ID: seeded.ID, + UpdatedAt: seeded.UpdatedAt, + OauthRefreshFailureReason: "stale", + }) + require.ErrorIs(t, err, sql.ErrNoRows) + //nolint:gocritic // Verifying persisted state requires system access. + row, err = db.GetMCPServerUserToken(dbauthz.AsSystemRestricted(ctx), database.GetMCPServerUserTokenParams{ + MCPServerConfigID: created.ID, + UserID: member.ID, + }) + require.NoError(t, err) + require.NotEqual(t, "stale", row.OauthRefreshFailureReason) + + // Re-authenticating (upserting fresh token material) clears the + // failure and restores connected status. + //nolint:gocritic // Seeding test state requires system access. + _, err = db.UpsertMCPServerUserToken(dbauthz.AsSystemRestricted(ctx), database.UpsertMCPServerUserTokenParams{ + MCPServerConfigID: created.ID, + UserID: member.ID, + AccessToken: "new-access", + RefreshToken: "new-refresh", + TokenType: "Bearer", + Expiry: sql.NullTime{Time: time.Now().Add(time.Hour), Valid: true}, + }) + require.NoError(t, err) + + configs, err = memberClient.MCPServerConfigs(ctx) + require.NoError(t, err) + require.Len(t, configs, 1) + require.True(t, configs[0].AuthConnected) + // The token is valid, so no refresh call is made. + require.Equal(t, hitsAfterFirstList, tokenEndpointHits.Load()) +} + +func TestMCPServerConfigsTransientRefreshFailure(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitLong) + providerKeys := coderdtest.FakeOpenAICompatProviderAPIKeys(t) + adminClient, db := coderdtest.NewWithDatabase(t, &coderdtest.Options{ + DeploymentValues: mcpDeploymentValues(t), + ChatProviderAPIKeys: &providerKeys, + }) + firstUser := coderdtest.CreateFirstUser(t, adminClient) + memberClient, member := coderdtest.CreateAnotherUser(t, adminClient, firstUser.OrganizationID) + + tokenSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusInternalServerError) + })) + t.Cleanup(tokenSrv.Close) + + created, err := adminClient.CreateMCPServerConfig(ctx, codersdk.CreateMCPServerConfigRequest{ + DisplayName: "Flaky Server", + Slug: "flaky-server", + Transport: "streamable_http", + URL: "https://mcp.example.com/v1", + AuthType: "oauth2", + OAuth2ClientID: "cid", + OAuth2AuthURL: "https://auth.example.com/authorize", + OAuth2TokenURL: tokenSrv.URL, + Availability: "default_on", + Enabled: true, + ToolAllowList: []string{}, + ToolDenyList: []string{}, + }) + require.NoError(t, err) + + //nolint:gocritic // Seeding test state requires system access. + _, err = db.UpsertMCPServerUserToken(dbauthz.AsSystemRestricted(ctx), database.UpsertMCPServerUserTokenParams{ + MCPServerConfigID: created.ID, + UserID: member.ID, + AccessToken: "expired-access", + RefreshToken: "still-good-refresh", + TokenType: "Bearer", + Expiry: sql.NullTime{Time: time.Now().Add(-time.Hour), Valid: true}, + }) + require.NoError(t, err) + + configs, err := memberClient.MCPServerConfigs(ctx) + require.NoError(t, err) + require.Len(t, configs, 1) + require.False(t, configs[0].AuthConnected) + + // Transient failures must not destroy the token: a later refresh + // may succeed. + //nolint:gocritic // Verifying persisted state requires system access. + row, err := db.GetMCPServerUserToken(dbauthz.AsSystemRestricted(ctx), database.GetMCPServerUserTokenParams{ + MCPServerConfigID: created.ID, + UserID: member.ID, + }) + require.NoError(t, err) + require.Equal(t, "still-good-refresh", row.RefreshToken) + require.Empty(t, row.OauthRefreshFailureReason) +} diff --git a/coderd/x/chatd/chatd.go b/coderd/x/chatd/chatd.go index 09bade41a8..bdf24fb71d 100644 --- a/coderd/x/chatd/chatd.go +++ b/coderd/x/chatd/chatd.go @@ -4747,6 +4747,12 @@ func (p *Server) refreshExpiredMCPTokens( if tok.RefreshToken == "" { continue } + if tok.OauthRefreshFailureReason != "" { + // A previous refresh already failed permanently (e.g. + // revoked grant); the user must reconnect. Skip the + // provider call entirely. + continue + } eg.Go(func() error { refreshed, err := p.refreshMCPTokenIfNeeded(ctx, logger, cfg, tok) @@ -4778,6 +4784,9 @@ func (p *Server) refreshMCPTokenIfNeeded( ) (database.MCPServerUserToken, error) { result, err := mcpclient.RefreshOAuth2Token(ctx, cfg, tok) if err != nil { + if mcpclient.IsPermanentRefreshError(err) { + return p.markMCPTokenRefreshFailure(ctx, logger, cfg, tok, err), nil + } return tok, err } @@ -4827,3 +4836,66 @@ func (p *Server) refreshMCPTokenIfNeeded( return updated, nil } + +// markMCPTokenRefreshFailure persists a permanent refresh failure +// (e.g. the upstream grant was revoked) so the dead token is never +// attached again and the UI can prompt the user to reconnect. It +// always returns a token with cleared auth material, even when +// persistence fails, so the current request does not send a stale +// bearer token. +func (p *Server) markMCPTokenRefreshFailure( + ctx context.Context, + logger slog.Logger, + cfg database.MCPServerConfig, + tok database.MCPServerUserToken, + refreshErr error, +) database.MCPServerUserToken { + logger.Warn(ctx, "mcp oauth2 grant permanently unusable, marking token for reconnect", + slog.F("server_slug", cfg.Slug), + slog.F("user_id", tok.UserID), + slog.Error(refreshErr), + ) + + //nolint:gocritic // Chatd needs system-level write access to + // persist the refresh failure for the user. + marked, err := p.db.MarkMCPServerUserTokenRefreshFailure( + dbauthz.AsSystemRestricted(ctx), + database.MarkMCPServerUserTokenRefreshFailureParams{ + ID: tok.ID, + UpdatedAt: tok.UpdatedAt, + OauthRefreshFailureReason: mcpclient.RefreshFailureReason(refreshErr), + }, + ) + if err == nil { + return marked + } + + if xerrors.Is(err, sql.ErrNoRows) { + // Optimistic lock miss: a concurrent request refreshed or + // replaced the token after we read it, so our failure is + // stale. Use the winner's row instead. + //nolint:gocritic // Chatd needs system-level read access to + // load the concurrently updated token. + current, readErr := p.db.GetMCPServerUserToken( + dbauthz.AsSystemRestricted(ctx), + database.GetMCPServerUserTokenParams{ + MCPServerConfigID: tok.MCPServerConfigID, + UserID: tok.UserID, + }, + ) + if readErr == nil { + return current + } + err = readErr + } + + logger.Warn(ctx, "failed to persist MCP oauth2 refresh failure", + slog.F("server_slug", cfg.Slug), + slog.Error(err), + ) + tok.AccessToken = "" + tok.RefreshToken = "" + tok.Expiry = sql.NullTime{} + tok.OauthRefreshFailureReason = mcpclient.RefreshFailureReason(refreshErr) + return tok +} diff --git a/coderd/x/chatd/mcp_refresh_internal_test.go b/coderd/x/chatd/mcp_refresh_internal_test.go new file mode 100644 index 0000000000..372fb6a7be --- /dev/null +++ b/coderd/x/chatd/mcp_refresh_internal_test.go @@ -0,0 +1,215 @@ +package chatd + +import ( + "context" + "database/sql" + "net/http" + "net/http/httptest" + "sync/atomic" + "testing" + "time" + + "github.com/google/uuid" + "github.com/stretchr/testify/require" + "go.uber.org/mock/gomock" + + "cdr.dev/slog/v3/sloggers/slogtest" + "github.com/coder/coder/v2/coderd/database" + "github.com/coder/coder/v2/coderd/database/dbmock" +) + +func invalidGrantServer(t *testing.T, hits *atomic.Int64) *httptest.Server { + t.Helper() + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + if hits != nil { + hits.Add(1) + } + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusBadRequest) + _, _ = w.Write([]byte(`{"error":"invalid_grant"}`)) + })) + t.Cleanup(srv.Close) + return srv +} + +func expiredMCPToken(cfgID uuid.UUID) database.MCPServerUserToken { + return database.MCPServerUserToken{ + ID: uuid.New(), + MCPServerConfigID: cfgID, + UserID: uuid.New(), + AccessToken: "expired-access", + RefreshToken: "dead-refresh", + TokenType: "Bearer", + Expiry: sql.NullTime{Time: time.Now().Add(-time.Hour), Valid: true}, + UpdatedAt: time.Now().Add(-time.Hour), + } +} + +func TestRefreshMCPTokenPermanentFailure(t *testing.T) { + t.Parallel() + + t.Run("MarksTokenAndClearsAuth", func(t *testing.T) { + t.Parallel() + tokenSrv := invalidGrantServer(t, nil) + + cfg := database.MCPServerConfig{ + ID: uuid.New(), + Slug: "revoked", + AuthType: "oauth2", + OAuth2ClientID: "cid", + OAuth2TokenURL: tokenSrv.URL, + } + tok := expiredMCPToken(cfg.ID) + + ctrl := gomock.NewController(t) + db := dbmock.NewMockStore(ctrl) + marked := tok + marked.AccessToken = "" + marked.RefreshToken = "" + marked.Expiry = sql.NullTime{} + marked.OauthRefreshFailureReason = "invalid_grant" + db.EXPECT(). + MarkMCPServerUserTokenRefreshFailure(gomock.Any(), gomock.Any()). + DoAndReturn(func(_ context.Context, arg database.MarkMCPServerUserTokenRefreshFailureParams) (database.MCPServerUserToken, error) { + require.Equal(t, tok.ID, arg.ID) + require.Equal(t, tok.UpdatedAt, arg.UpdatedAt) + require.Contains(t, arg.OauthRefreshFailureReason, "invalid_grant") + return marked, nil + }) + + server := &Server{db: db} + result, err := server.refreshMCPTokenIfNeeded( + context.Background(), slogtest.Make(t, nil), cfg, tok, + ) + require.NoError(t, err) + require.Empty(t, result.AccessToken) + require.Empty(t, result.RefreshToken) + require.NotEmpty(t, result.OauthRefreshFailureReason) + }) + + t.Run("OptimisticLockLossUsesWinnerRow", func(t *testing.T) { + t.Parallel() + tokenSrv := invalidGrantServer(t, nil) + + cfg := database.MCPServerConfig{ + ID: uuid.New(), + Slug: "raced", + AuthType: "oauth2", + OAuth2ClientID: "cid", + OAuth2TokenURL: tokenSrv.URL, + } + tok := expiredMCPToken(cfg.ID) + winner := tok + winner.AccessToken = "fresh-access" + winner.RefreshToken = "fresh-refresh" + winner.Expiry = sql.NullTime{Time: time.Now().Add(time.Hour), Valid: true} + + ctrl := gomock.NewController(t) + db := dbmock.NewMockStore(ctrl) + db.EXPECT(). + MarkMCPServerUserTokenRefreshFailure(gomock.Any(), gomock.Any()). + Return(database.MCPServerUserToken{}, sql.ErrNoRows) + db.EXPECT(). + GetMCPServerUserToken(gomock.Any(), database.GetMCPServerUserTokenParams{ + MCPServerConfigID: tok.MCPServerConfigID, + UserID: tok.UserID, + }). + Return(winner, nil) + + server := &Server{db: db} + result, err := server.refreshMCPTokenIfNeeded( + context.Background(), slogtest.Make(t, nil), cfg, tok, + ) + require.NoError(t, err) + require.Equal(t, "fresh-access", result.AccessToken) + require.Empty(t, result.OauthRefreshFailureReason) + }) + + t.Run("PersistFailureStillClearsAuth", func(t *testing.T) { + t.Parallel() + tokenSrv := invalidGrantServer(t, nil) + + cfg := database.MCPServerConfig{ + ID: uuid.New(), + Slug: "db-down", + AuthType: "oauth2", + OAuth2ClientID: "cid", + OAuth2TokenURL: tokenSrv.URL, + } + tok := expiredMCPToken(cfg.ID) + + ctrl := gomock.NewController(t) + db := dbmock.NewMockStore(ctrl) + db.EXPECT(). + MarkMCPServerUserTokenRefreshFailure(gomock.Any(), gomock.Any()). + Return(database.MCPServerUserToken{}, sql.ErrConnDone) + + server := &Server{db: db} + result, err := server.refreshMCPTokenIfNeeded( + context.Background(), slogtest.Make(t, nil), cfg, tok, + ) + require.NoError(t, err) + require.Empty(t, result.AccessToken) + require.Empty(t, result.RefreshToken) + require.NotEmpty(t, result.OauthRefreshFailureReason) + }) + + t.Run("TransientFailureKeepsToken", func(t *testing.T) { + t.Parallel() + tokenSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusInternalServerError) + })) + t.Cleanup(tokenSrv.Close) + + cfg := database.MCPServerConfig{ + ID: uuid.New(), + Slug: "flaky", + AuthType: "oauth2", + OAuth2ClientID: "cid", + OAuth2TokenURL: tokenSrv.URL, + } + tok := expiredMCPToken(cfg.ID) + + ctrl := gomock.NewController(t) + db := dbmock.NewMockStore(ctrl) + + server := &Server{db: db} + result, err := server.refreshMCPTokenIfNeeded( + context.Background(), slogtest.Make(t, nil), cfg, tok, + ) + require.Error(t, err) + require.Equal(t, tok.AccessToken, result.AccessToken) + require.Equal(t, tok.RefreshToken, result.RefreshToken) + require.Empty(t, result.OauthRefreshFailureReason) + }) +} + +func TestRefreshExpiredMCPTokensSkipsFailedTokens(t *testing.T) { + t.Parallel() + + var hits atomic.Int64 + tokenSrv := invalidGrantServer(t, &hits) + + cfg := database.MCPServerConfig{ + ID: uuid.New(), + Slug: "revoked", + AuthType: "oauth2", + OAuth2ClientID: "cid", + OAuth2TokenURL: tokenSrv.URL, + } + tok := expiredMCPToken(cfg.ID) + tok.OauthRefreshFailureReason = "invalid_grant" + + ctrl := gomock.NewController(t) + db := dbmock.NewMockStore(ctrl) + + server := &Server{db: db} + result := server.refreshExpiredMCPTokens( + context.Background(), slogtest.Make(t, nil), + []database.MCPServerConfig{cfg}, + []database.MCPServerUserToken{tok}, + ) + require.Len(t, result, 1) + require.Equal(t, tok, result[0]) + require.EqualValues(t, 0, hits.Load(), "provider must not be called for failed tokens") +} diff --git a/coderd/x/chatd/mcpclient/export_test.go b/coderd/x/chatd/mcpclient/export_test.go index dbdca8c638..50d350aba2 100644 --- a/coderd/x/chatd/mcpclient/export_test.go +++ b/coderd/x/chatd/mcpclient/export_test.go @@ -3,3 +3,7 @@ package mcpclient // ConvertCallResultForTest exposes convertCallResult for external // tests. var ConvertCallResultForTest = convertCallResult + +// BuildAuthHeadersForTest exposes buildAuthHeaders for external +// tests. +var BuildAuthHeadersForTest = buildAuthHeaders diff --git a/coderd/x/chatd/mcpclient/mcpclient.go b/coderd/x/chatd/mcpclient/mcpclient.go index 0ff65db0aa..4214ba42c4 100644 --- a/coderd/x/chatd/mcpclient/mcpclient.go +++ b/coderd/x/chatd/mcpclient/mcpclient.go @@ -370,6 +370,16 @@ func buildAuthHeaders( ) break } + if tok.OauthRefreshFailureReason != "" { + // The grant is permanently unusable (e.g. revoked + // upstream) and the user must reconnect. Do not attach + // any leftover token material. + logger.Warn(ctx, + "oauth2 token for MCP server requires reconnect, skipping auth header", + slog.F("server_slug", cfg.Slug), + ) + break + } if tok.Expiry.Valid && tok.Expiry.Time.Before(time.Now()) { logger.Warn(ctx, "oauth2 token for MCP server is expired", @@ -858,6 +868,43 @@ type RefreshResult struct { Refreshed bool } +// refreshFailureReasonLimit caps the error text persisted to +// mcp_server_user_tokens.oauth_refresh_failure_reason, matching the +// external auth failure reason limit. +const refreshFailureReasonLimit = 400 + +// IsPermanentRefreshError reports whether an OAuth2 token refresh +// error means the user's grant is permanently unusable (for example +// the upstream grant was revoked) rather than a transient provider +// failure. Only error codes tied to the grant itself count: client +// or config problems (invalid_client, unauthorized_client, ...) +// affect every user of the server and cannot be fixed by the user +// reconnecting, so they are treated as transient here. See RFC 6749 +// section 5.2. +func IsPermanentRefreshError(err error) bool { + var oauthErr *oauth2.RetrieveError + if !xerrors.As(err, &oauthErr) { + return false + } + switch oauthErr.ErrorCode { + case "invalid_grant", // RFC 6749: grant invalid, expired, or revoked + "bad_refresh_token": // GitHub's equivalent, returned with HTTP 200 + return true + } + return false +} + +// RefreshFailureReason converts a refresh error into a bounded string +// safe to persist as oauth_refresh_failure_reason. It is stored for +// operator debugging only and is never returned through the API. +func RefreshFailureReason(err error) string { + reason := err.Error() + if len(reason) > refreshFailureReasonLimit { + reason = reason[:refreshFailureReasonLimit] + } + return reason +} + // RefreshOAuth2Token checks whether the given MCP user token is // expired (or within 10 seconds of expiry) and refreshes it using // the OAuth2 credentials from the server config. If the token is diff --git a/coderd/x/chatd/mcpclient/refresh_test.go b/coderd/x/chatd/mcpclient/refresh_test.go new file mode 100644 index 0000000000..901cba432f --- /dev/null +++ b/coderd/x/chatd/mcpclient/refresh_test.go @@ -0,0 +1,136 @@ +package mcpclient_test + +import ( + "context" + "database/sql" + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/google/uuid" + "github.com/stretchr/testify/require" + "golang.org/x/oauth2" + "golang.org/x/xerrors" + + "cdr.dev/slog/v3/sloggers/slogtest" + "github.com/coder/coder/v2/coderd/database" + "github.com/coder/coder/v2/coderd/x/chatd/mcpclient" +) + +func TestIsPermanentRefreshError(t *testing.T) { + t.Parallel() + + retrieveErr := func(code string, status int) error { + body, err := json.Marshal(map[string]string{"error": code}) + require.NoError(t, err) + return &oauth2.RetrieveError{ + Response: &http.Response{StatusCode: status}, + Body: body, + ErrorCode: code, + } + } + + cases := []struct { + name string + err error + permanent bool + }{ + {"InvalidGrant", retrieveErr("invalid_grant", http.StatusBadRequest), true}, + {"BadRefreshToken", retrieveErr("bad_refresh_token", http.StatusOK), true}, + {"WrappedInvalidGrant", xerrors.Errorf("refresh: %w", retrieveErr("invalid_grant", http.StatusBadRequest)), true}, + {"InvalidClient", retrieveErr("invalid_client", http.StatusUnauthorized), false}, + {"UnauthorizedClient", retrieveErr("unauthorized_client", http.StatusBadRequest), false}, + {"ServerError", retrieveErr("", http.StatusInternalServerError), false}, + {"RateLimited", retrieveErr("", http.StatusTooManyRequests), false}, + {"PlainError", xerrors.New("connection refused"), false}, + {"Nil", nil, false}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + require.Equal(t, tc.permanent, mcpclient.IsPermanentRefreshError(tc.err)) + }) + } +} + +func TestRefreshFailureReason(t *testing.T) { + t.Parallel() + + require.Equal(t, "boom", mcpclient.RefreshFailureReason(xerrors.New("boom"))) + + long := strings.Repeat("x", 1000) + reason := mcpclient.RefreshFailureReason(xerrors.New(long)) + require.Len(t, reason, 400) +} + +func TestRefreshOAuth2TokenInvalidGrant(t *testing.T) { + t.Parallel() + + tokenSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusBadRequest) + _, _ = w.Write([]byte(`{"error":"invalid_grant","error_description":"grant revoked"}`)) + })) + defer tokenSrv.Close() + + cfg := database.MCPServerConfig{ + OAuth2ClientID: "cid", + OAuth2TokenURL: tokenSrv.URL, + } + tok := database.MCPServerUserToken{ + AccessToken: "expired", + RefreshToken: "refresh", + TokenType: "Bearer", + Expiry: sql.NullTime{Time: time.Now().Add(-time.Hour), Valid: true}, + } + + _, err := mcpclient.RefreshOAuth2Token(context.Background(), cfg, tok) + require.Error(t, err) + require.True(t, mcpclient.IsPermanentRefreshError(err)) +} + +func TestBuildAuthHeadersSkipsFailedToken(t *testing.T) { + t.Parallel() + + logger := slogtest.Make(t, nil) + cfg := database.MCPServerConfig{ + ID: uuid.New(), + Slug: "revoked", + AuthType: "oauth2", + } + + t.Run("FailureReasonSet", func(t *testing.T) { + t.Parallel() + headers := mcpclient.BuildAuthHeadersForTest( + context.Background(), logger, cfg, + map[uuid.UUID]database.MCPServerUserToken{ + cfg.ID: { + MCPServerConfigID: cfg.ID, + AccessToken: "leftover", + OauthRefreshFailureReason: "invalid_grant", + }, + }, + uuid.New(), nil, + ) + require.NotContains(t, headers, "Authorization") + }) + + t.Run("HealthyToken", func(t *testing.T) { + t.Parallel() + headers := mcpclient.BuildAuthHeadersForTest( + context.Background(), logger, cfg, + map[uuid.UUID]database.MCPServerUserToken{ + cfg.ID: { + MCPServerConfigID: cfg.ID, + AccessToken: "valid", + TokenType: "Bearer", + }, + }, + uuid.New(), nil, + ) + require.Equal(t, "Bearer valid", headers["Authorization"]) + }) +} diff --git a/enterprise/dbcrypt/dbcrypt.go b/enterprise/dbcrypt/dbcrypt.go index d8998baddf..de6211f2fb 100644 --- a/enterprise/dbcrypt/dbcrypt.go +++ b/enterprise/dbcrypt/dbcrypt.go @@ -787,6 +787,20 @@ func (db *dbCrypt) GetMCPServerUserToken(ctx context.Context, arg database.GetMC return tok, nil } +func (db *dbCrypt) MarkMCPServerUserTokenRefreshFailure(ctx context.Context, params database.MarkMCPServerUserTokenRefreshFailureParams) (database.MCPServerUserToken, error) { + // The query clears the encrypted token fields, so nothing needs + // encrypting; decrypt the returned row for consistency with the + // other accessors (a no-op for the cleared fields). + tok, err := db.Store.MarkMCPServerUserTokenRefreshFailure(ctx, params) + if err != nil { + return database.MCPServerUserToken{}, err + } + if err := db.decryptMCPServerUserToken(&tok); err != nil { + return database.MCPServerUserToken{}, err + } + return tok, nil +} + func (db *dbCrypt) GetMCPServerUserTokensByUserID(ctx context.Context, userID uuid.UUID) ([]database.MCPServerUserToken, error) { toks, err := db.Store.GetMCPServerUserTokensByUserID(ctx, userID) if err != nil { diff --git a/enterprise/dbcrypt/dbcrypt_internal_test.go b/enterprise/dbcrypt/dbcrypt_internal_test.go index bcfd9d41da..a42fb221e0 100644 --- a/enterprise/dbcrypt/dbcrypt_internal_test.go +++ b/enterprise/dbcrypt/dbcrypt_internal_test.go @@ -1572,6 +1572,24 @@ func TestMCPServerUserTokens(t *testing.T) { requireEncryptedEquals(t, ciphers[0], rawTok.AccessToken, accessToken) requireEncryptedEquals(t, ciphers[0], rawTok.RefreshToken, refreshToken) }) + + t.Run("MarkMCPServerUserTokenRefreshFailure", func(t *testing.T) { + t.Parallel() + _, crypt, ciphers := setup(t) + _, tok := insertConfigAndToken(t, crypt, ciphers) + + marked, err := crypt.MarkMCPServerUserTokenRefreshFailure(ctx, database.MarkMCPServerUserTokenRefreshFailureParams{ + ID: tok.ID, + UpdatedAt: tok.UpdatedAt, + OauthRefreshFailureReason: "invalid_grant", + }) + require.NoError(t, err) + require.Empty(t, marked.AccessToken) + require.Empty(t, marked.RefreshToken) + require.False(t, marked.AccessTokenKeyID.Valid) + require.False(t, marked.RefreshTokenKeyID.Valid) + require.Equal(t, "invalid_grant", marked.OauthRefreshFailureReason) + }) } func TestUserSecrets(t *testing.T) {