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.
This commit is contained in:
Michael Suchacz
2026-07-16 11:43:05 +00:00
committed by GitHub
parent 213f5ce606
commit e489092154
20 changed files with 913 additions and 25 deletions
+72
View File
@@ -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
}
+215
View File
@@ -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")
}
+4
View File
@@ -3,3 +3,7 @@ package mcpclient
// ConvertCallResultForTest exposes convertCallResult for external
// tests.
var ConvertCallResultForTest = convertCallResult
// BuildAuthHeadersForTest exposes buildAuthHeaders for external
// tests.
var BuildAuthHeadersForTest = buildAuthHeaders
+47
View File
@@ -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
+136
View File
@@ -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"])
})
}