mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
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:
@@ -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
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
@@ -3,3 +3,7 @@ package mcpclient
|
||||
// ConvertCallResultForTest exposes convertCallResult for external
|
||||
// tests.
|
||||
var ConvertCallResultForTest = convertCallResult
|
||||
|
||||
// BuildAuthHeadersForTest exposes buildAuthHeaders for external
|
||||
// tests.
|
||||
var BuildAuthHeadersForTest = buildAuthHeaders
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"])
|
||||
})
|
||||
}
|
||||
Reference in New Issue
Block a user