mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat: revoke MCP server OAuth grants at the provider on disconnect (#27300)
Closes [CODAGT-805](https://linear.app/codercom/issue/CODAGT-805/revoke-oauth-grants-at-the-source-for-mcp-servers). The experimental MCP server OAuth2 disconnect endpoint previously deleted only the stored token row, leaving the grant active at the OAuth provider. This PR adds provider-side token revocation while keeping local disconnect independent of provider availability. ## Changes - Add `mcp_server_configs.oauth2_revocation_url` in migration `000547`. The value can be configured manually, discovered from RFC 8414 metadata, and managed through the MCP server settings UI. Non-admin responses redact it with the other OAuth2 fields. - Revoke the refresh token first through the RFC 7009 endpoint, then fall back to the access token only for `unsupported_token_type`. Public clients send `client_id`; confidential clients use `client_secret_basic`. - Delete the local token transactionally before best-effort provider revocation. Callers without a token receive the same response for hidden and nonexistent config IDs, and provider failures return a generic warning without exposing provider response bodies. - Require HTTPS revocation endpoints except for HTTP loopback URLs. Redirects must preserve the POST and remain on the configured origin. Redirect errors omit provider-controlled paths and query strings so reflected token material cannot enter logs. - Treat `200 OK` and `204 No Content` as completed revocations. `202 Accepted` remains a failure because it does not confirm completion. - Prevent an in-flight refresh from recreating a token deleted by disconnect. Refresh persistence now uses an optimistic update keyed by token ID and `updated_at`; only the OAuth callback can create a token row. Refresh conflicts reload the current row or clear in-memory auth when disconnect deleted it. - Return `{token_revoked, token_revocation_error}` from disconnect, while retaining SDK compatibility with the legacy `204` response. The UI surfaces provider revocation failures as warning toasts. - Document revocation endpoint discovery, HTTPS requirements, and best-effort disconnect behavior. No token or no configured revocation URL returns `token_revoked: false` without an error, so disconnect remains idempotent. > Updated by Mux, an AI coding agent, on Mike's behalf.
This commit is contained in:
@@ -7599,6 +7599,13 @@ func (q *querier) UpdateMCPServerConfig(ctx context.Context, arg database.Update
|
||||
return q.db.UpdateMCPServerConfig(ctx, arg)
|
||||
}
|
||||
|
||||
func (q *querier) UpdateMCPServerUserTokenFromRefresh(ctx context.Context, arg database.UpdateMCPServerUserTokenFromRefreshParams) (database.MCPServerUserToken, error) {
|
||||
if err := q.authorizeContext(ctx, policy.ActionUpdate, rbac.ResourceDeploymentConfig); err != nil {
|
||||
return database.MCPServerUserToken{}, err
|
||||
}
|
||||
return q.db.UpdateMCPServerUserTokenFromRefresh(ctx, arg)
|
||||
}
|
||||
|
||||
func (q *querier) UpdateMemberRoles(ctx context.Context, arg database.UpdateMemberRolesParams) (database.OrganizationMember, error) {
|
||||
// Authorized fetch will check that the actor has read access to the org member since the org member is returned.
|
||||
member, err := database.ExpectOne(q.OrganizationMembers(ctx, database.OrganizationMembersParams{
|
||||
|
||||
@@ -1937,6 +1937,17 @@ func (s *MethodTestSuite) TestChats() {
|
||||
dbm.EXPECT().UpdateMCPServerConfig(gomock.Any(), arg).Return(config, nil).AnyTimes()
|
||||
check.Args(arg).Asserts(rbac.ResourceDeploymentConfig, policy.ActionUpdate).Returns(config)
|
||||
}))
|
||||
s.Run("UpdateMCPServerUserTokenFromRefresh", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
|
||||
token := testutil.Fake(s.T(), faker, database.MCPServerUserToken{})
|
||||
arg := database.UpdateMCPServerUserTokenFromRefreshParams{
|
||||
ID: token.ID,
|
||||
UpdatedAt: token.UpdatedAt,
|
||||
AccessToken: "refreshed-access-token",
|
||||
TokenType: "bearer",
|
||||
}
|
||||
dbm.EXPECT().UpdateMCPServerUserTokenFromRefresh(gomock.Any(), arg).Return(token, nil).AnyTimes()
|
||||
check.Args(arg).Asserts(rbac.ResourceDeploymentConfig, policy.ActionUpdate).Returns(token)
|
||||
}))
|
||||
s.Run("UpsertMCPServerUserToken", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
|
||||
arg := database.UpsertMCPServerUserTokenParams{
|
||||
MCPServerConfigID: uuid.New(),
|
||||
|
||||
@@ -361,6 +361,7 @@ func MCPServerConfig(t testing.TB, db database.Store, seed database.MCPServerCon
|
||||
OAuth2ClientSecretKeyID: seed.OAuth2ClientSecretKeyID,
|
||||
OAuth2AuthURL: seed.OAuth2AuthURL,
|
||||
OAuth2TokenURL: seed.OAuth2TokenURL,
|
||||
OAuth2RevocationURL: seed.OAuth2RevocationURL,
|
||||
OAuth2Scopes: seed.OAuth2Scopes,
|
||||
APIKeyHeader: seed.APIKeyHeader,
|
||||
APIKeyValue: seed.APIKeyValue,
|
||||
|
||||
+8
@@ -5393,6 +5393,14 @@ func (m queryMetricsStore) UpdateMCPServerConfig(ctx context.Context, arg databa
|
||||
return r0, r1
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) UpdateMCPServerUserTokenFromRefresh(ctx context.Context, arg database.UpdateMCPServerUserTokenFromRefreshParams) (database.MCPServerUserToken, error) {
|
||||
start := time.Now()
|
||||
r0, r1 := m.s.UpdateMCPServerUserTokenFromRefresh(ctx, arg)
|
||||
m.queryLatencies.WithLabelValues("UpdateMCPServerUserTokenFromRefresh").Observe(time.Since(start).Seconds())
|
||||
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "UpdateMCPServerUserTokenFromRefresh").Inc()
|
||||
return r0, r1
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) UpdateMemberRoles(ctx context.Context, arg database.UpdateMemberRolesParams) (database.OrganizationMember, error) {
|
||||
start := time.Now()
|
||||
r0, r1 := m.s.UpdateMemberRoles(ctx, arg)
|
||||
|
||||
Generated
+15
@@ -10159,6 +10159,21 @@ func (mr *MockStoreMockRecorder) UpdateMCPServerConfig(ctx, arg any) *gomock.Cal
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpdateMCPServerConfig", reflect.TypeOf((*MockStore)(nil).UpdateMCPServerConfig), ctx, arg)
|
||||
}
|
||||
|
||||
// UpdateMCPServerUserTokenFromRefresh mocks base method.
|
||||
func (m *MockStore) UpdateMCPServerUserTokenFromRefresh(ctx context.Context, arg database.UpdateMCPServerUserTokenFromRefreshParams) (database.MCPServerUserToken, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "UpdateMCPServerUserTokenFromRefresh", ctx, arg)
|
||||
ret0, _ := ret[0].(database.MCPServerUserToken)
|
||||
ret1, _ := ret[1].(error)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// UpdateMCPServerUserTokenFromRefresh indicates an expected call of UpdateMCPServerUserTokenFromRefresh.
|
||||
func (mr *MockStoreMockRecorder) UpdateMCPServerUserTokenFromRefresh(ctx, arg any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpdateMCPServerUserTokenFromRefresh", reflect.TypeOf((*MockStore)(nil).UpdateMCPServerUserTokenFromRefresh), ctx, arg)
|
||||
}
|
||||
|
||||
// UpdateMemberRoles mocks base method.
|
||||
func (m *MockStore) UpdateMemberRoles(ctx context.Context, arg database.UpdateMemberRolesParams) (database.OrganizationMember, error) {
|
||||
m.ctrl.T.Helper()
|
||||
|
||||
Generated
+1
@@ -2496,6 +2496,7 @@ CREATE TABLE mcp_server_configs (
|
||||
model_intent boolean DEFAULT false NOT NULL,
|
||||
allow_in_plan_mode boolean DEFAULT false NOT NULL,
|
||||
forward_coder_headers boolean DEFAULT false NOT NULL,
|
||||
oauth2_revocation_url text DEFAULT ''::text NOT NULL,
|
||||
CONSTRAINT mcp_server_configs_auth_type_check CHECK ((auth_type = ANY (ARRAY['none'::text, 'oauth2'::text, 'api_key'::text, 'custom_headers'::text, 'user_oidc'::text]))),
|
||||
CONSTRAINT mcp_server_configs_availability_check CHECK ((availability = ANY (ARRAY['force_on'::text, 'default_on'::text, 'default_off'::text]))),
|
||||
CONSTRAINT mcp_server_configs_transport_check CHECK ((transport = ANY (ARRAY['streamable_http'::text, 'sse'::text])))
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
ALTER TABLE mcp_server_configs
|
||||
DROP COLUMN oauth2_revocation_url;
|
||||
@@ -0,0 +1,2 @@
|
||||
ALTER TABLE mcp_server_configs
|
||||
ADD COLUMN oauth2_revocation_url text NOT NULL DEFAULT '';
|
||||
Generated
+1
@@ -5437,6 +5437,7 @@ type MCPServerConfig struct {
|
||||
ModelIntent bool `db:"model_intent" json:"model_intent"`
|
||||
AllowInPlanMode bool `db:"allow_in_plan_mode" json:"allow_in_plan_mode"`
|
||||
ForwardCoderHeaders bool `db:"forward_coder_headers" json:"forward_coder_headers"`
|
||||
OAuth2RevocationURL string `db:"oauth2_revocation_url" json:"oauth2_revocation_url"`
|
||||
}
|
||||
|
||||
type MCPServerUserToken struct {
|
||||
|
||||
Generated
+3
@@ -1428,6 +1428,9 @@ type sqlcQuerier interface {
|
||||
UpdateInactiveUsersToDormant(ctx context.Context, arg UpdateInactiveUsersToDormantParams) ([]UpdateInactiveUsersToDormantRow, error)
|
||||
UpdateInboxNotificationReadStatus(ctx context.Context, arg UpdateInboxNotificationReadStatusParams) error
|
||||
UpdateMCPServerConfig(ctx context.Context, arg UpdateMCPServerConfigParams) (MCPServerConfig, error)
|
||||
// Refresh persistence must not recreate a token deleted by disconnect.
|
||||
// The optimistic lock also prevents stale refreshes from replacing newer tokens.
|
||||
UpdateMCPServerUserTokenFromRefresh(ctx context.Context, arg UpdateMCPServerUserTokenFromRefreshParams) (MCPServerUserToken, error)
|
||||
UpdateMemberRoles(ctx context.Context, arg UpdateMemberRolesParams) (OrganizationMember, error)
|
||||
UpdateMemoryResourceMonitor(ctx context.Context, arg UpdateMemoryResourceMonitorParams) error
|
||||
UpdateNotificationTemplateMethodByID(ctx context.Context, arg UpdateNotificationTemplateMethodByIDParams) (NotificationTemplate, error)
|
||||
|
||||
Generated
+103
-28
@@ -16693,7 +16693,7 @@ func (q *sqlQuerier) DeleteMCPServerUserToken(ctx context.Context, arg DeleteMCP
|
||||
|
||||
const getEnabledMCPServerConfigs = `-- name: GetEnabledMCPServerConfigs :many
|
||||
SELECT
|
||||
id, display_name, slug, description, icon_url, transport, url, auth_type, oauth2_client_id, oauth2_client_secret, oauth2_client_secret_key_id, oauth2_auth_url, oauth2_token_url, oauth2_scopes, api_key_header, api_key_value, api_key_value_key_id, custom_headers, custom_headers_key_id, tool_allow_list, tool_deny_list, availability, enabled, created_by, updated_by, created_at, updated_at, model_intent, allow_in_plan_mode, forward_coder_headers
|
||||
id, display_name, slug, description, icon_url, transport, url, auth_type, oauth2_client_id, oauth2_client_secret, oauth2_client_secret_key_id, oauth2_auth_url, oauth2_token_url, oauth2_scopes, api_key_header, api_key_value, api_key_value_key_id, custom_headers, custom_headers_key_id, tool_allow_list, tool_deny_list, availability, enabled, created_by, updated_by, created_at, updated_at, model_intent, allow_in_plan_mode, forward_coder_headers, oauth2_revocation_url
|
||||
FROM
|
||||
mcp_server_configs
|
||||
WHERE
|
||||
@@ -16742,6 +16742,7 @@ func (q *sqlQuerier) GetEnabledMCPServerConfigs(ctx context.Context) ([]MCPServe
|
||||
&i.ModelIntent,
|
||||
&i.AllowInPlanMode,
|
||||
&i.ForwardCoderHeaders,
|
||||
&i.OAuth2RevocationURL,
|
||||
); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -16758,7 +16759,7 @@ func (q *sqlQuerier) GetEnabledMCPServerConfigs(ctx context.Context) ([]MCPServe
|
||||
|
||||
const getForcedMCPServerConfigs = `-- name: GetForcedMCPServerConfigs :many
|
||||
SELECT
|
||||
id, display_name, slug, description, icon_url, transport, url, auth_type, oauth2_client_id, oauth2_client_secret, oauth2_client_secret_key_id, oauth2_auth_url, oauth2_token_url, oauth2_scopes, api_key_header, api_key_value, api_key_value_key_id, custom_headers, custom_headers_key_id, tool_allow_list, tool_deny_list, availability, enabled, created_by, updated_by, created_at, updated_at, model_intent, allow_in_plan_mode, forward_coder_headers
|
||||
id, display_name, slug, description, icon_url, transport, url, auth_type, oauth2_client_id, oauth2_client_secret, oauth2_client_secret_key_id, oauth2_auth_url, oauth2_token_url, oauth2_scopes, api_key_header, api_key_value, api_key_value_key_id, custom_headers, custom_headers_key_id, tool_allow_list, tool_deny_list, availability, enabled, created_by, updated_by, created_at, updated_at, model_intent, allow_in_plan_mode, forward_coder_headers, oauth2_revocation_url
|
||||
FROM
|
||||
mcp_server_configs
|
||||
WHERE
|
||||
@@ -16808,6 +16809,7 @@ func (q *sqlQuerier) GetForcedMCPServerConfigs(ctx context.Context) ([]MCPServer
|
||||
&i.ModelIntent,
|
||||
&i.AllowInPlanMode,
|
||||
&i.ForwardCoderHeaders,
|
||||
&i.OAuth2RevocationURL,
|
||||
); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -16824,7 +16826,7 @@ func (q *sqlQuerier) GetForcedMCPServerConfigs(ctx context.Context) ([]MCPServer
|
||||
|
||||
const getMCPServerConfigByID = `-- name: GetMCPServerConfigByID :one
|
||||
SELECT
|
||||
id, display_name, slug, description, icon_url, transport, url, auth_type, oauth2_client_id, oauth2_client_secret, oauth2_client_secret_key_id, oauth2_auth_url, oauth2_token_url, oauth2_scopes, api_key_header, api_key_value, api_key_value_key_id, custom_headers, custom_headers_key_id, tool_allow_list, tool_deny_list, availability, enabled, created_by, updated_by, created_at, updated_at, model_intent, allow_in_plan_mode, forward_coder_headers
|
||||
id, display_name, slug, description, icon_url, transport, url, auth_type, oauth2_client_id, oauth2_client_secret, oauth2_client_secret_key_id, oauth2_auth_url, oauth2_token_url, oauth2_scopes, api_key_header, api_key_value, api_key_value_key_id, custom_headers, custom_headers_key_id, tool_allow_list, tool_deny_list, availability, enabled, created_by, updated_by, created_at, updated_at, model_intent, allow_in_plan_mode, forward_coder_headers, oauth2_revocation_url
|
||||
FROM
|
||||
mcp_server_configs
|
||||
WHERE
|
||||
@@ -16865,13 +16867,14 @@ func (q *sqlQuerier) GetMCPServerConfigByID(ctx context.Context, id uuid.UUID) (
|
||||
&i.ModelIntent,
|
||||
&i.AllowInPlanMode,
|
||||
&i.ForwardCoderHeaders,
|
||||
&i.OAuth2RevocationURL,
|
||||
)
|
||||
return i, err
|
||||
}
|
||||
|
||||
const getMCPServerConfigBySlug = `-- name: GetMCPServerConfigBySlug :one
|
||||
SELECT
|
||||
id, display_name, slug, description, icon_url, transport, url, auth_type, oauth2_client_id, oauth2_client_secret, oauth2_client_secret_key_id, oauth2_auth_url, oauth2_token_url, oauth2_scopes, api_key_header, api_key_value, api_key_value_key_id, custom_headers, custom_headers_key_id, tool_allow_list, tool_deny_list, availability, enabled, created_by, updated_by, created_at, updated_at, model_intent, allow_in_plan_mode, forward_coder_headers
|
||||
id, display_name, slug, description, icon_url, transport, url, auth_type, oauth2_client_id, oauth2_client_secret, oauth2_client_secret_key_id, oauth2_auth_url, oauth2_token_url, oauth2_scopes, api_key_header, api_key_value, api_key_value_key_id, custom_headers, custom_headers_key_id, tool_allow_list, tool_deny_list, availability, enabled, created_by, updated_by, created_at, updated_at, model_intent, allow_in_plan_mode, forward_coder_headers, oauth2_revocation_url
|
||||
FROM
|
||||
mcp_server_configs
|
||||
WHERE
|
||||
@@ -16912,13 +16915,14 @@ func (q *sqlQuerier) GetMCPServerConfigBySlug(ctx context.Context, slug string)
|
||||
&i.ModelIntent,
|
||||
&i.AllowInPlanMode,
|
||||
&i.ForwardCoderHeaders,
|
||||
&i.OAuth2RevocationURL,
|
||||
)
|
||||
return i, err
|
||||
}
|
||||
|
||||
const getMCPServerConfigs = `-- name: GetMCPServerConfigs :many
|
||||
SELECT
|
||||
id, display_name, slug, description, icon_url, transport, url, auth_type, oauth2_client_id, oauth2_client_secret, oauth2_client_secret_key_id, oauth2_auth_url, oauth2_token_url, oauth2_scopes, api_key_header, api_key_value, api_key_value_key_id, custom_headers, custom_headers_key_id, tool_allow_list, tool_deny_list, availability, enabled, created_by, updated_by, created_at, updated_at, model_intent, allow_in_plan_mode, forward_coder_headers
|
||||
id, display_name, slug, description, icon_url, transport, url, auth_type, oauth2_client_id, oauth2_client_secret, oauth2_client_secret_key_id, oauth2_auth_url, oauth2_token_url, oauth2_scopes, api_key_header, api_key_value, api_key_value_key_id, custom_headers, custom_headers_key_id, tool_allow_list, tool_deny_list, availability, enabled, created_by, updated_by, created_at, updated_at, model_intent, allow_in_plan_mode, forward_coder_headers, oauth2_revocation_url
|
||||
FROM
|
||||
mcp_server_configs
|
||||
ORDER BY
|
||||
@@ -16965,6 +16969,7 @@ func (q *sqlQuerier) GetMCPServerConfigs(ctx context.Context) ([]MCPServerConfig
|
||||
&i.ModelIntent,
|
||||
&i.AllowInPlanMode,
|
||||
&i.ForwardCoderHeaders,
|
||||
&i.OAuth2RevocationURL,
|
||||
); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -16981,7 +16986,7 @@ func (q *sqlQuerier) GetMCPServerConfigs(ctx context.Context) ([]MCPServerConfig
|
||||
|
||||
const getMCPServerConfigsByIDs = `-- name: GetMCPServerConfigsByIDs :many
|
||||
SELECT
|
||||
id, display_name, slug, description, icon_url, transport, url, auth_type, oauth2_client_id, oauth2_client_secret, oauth2_client_secret_key_id, oauth2_auth_url, oauth2_token_url, oauth2_scopes, api_key_header, api_key_value, api_key_value_key_id, custom_headers, custom_headers_key_id, tool_allow_list, tool_deny_list, availability, enabled, created_by, updated_by, created_at, updated_at, model_intent, allow_in_plan_mode, forward_coder_headers
|
||||
id, display_name, slug, description, icon_url, transport, url, auth_type, oauth2_client_id, oauth2_client_secret, oauth2_client_secret_key_id, oauth2_auth_url, oauth2_token_url, oauth2_scopes, api_key_header, api_key_value, api_key_value_key_id, custom_headers, custom_headers_key_id, tool_allow_list, tool_deny_list, availability, enabled, created_by, updated_by, created_at, updated_at, model_intent, allow_in_plan_mode, forward_coder_headers, oauth2_revocation_url
|
||||
FROM
|
||||
mcp_server_configs
|
||||
WHERE
|
||||
@@ -17030,6 +17035,7 @@ func (q *sqlQuerier) GetMCPServerConfigsByIDs(ctx context.Context, ids []uuid.UU
|
||||
&i.ModelIntent,
|
||||
&i.AllowInPlanMode,
|
||||
&i.ForwardCoderHeaders,
|
||||
&i.OAuth2RevocationURL,
|
||||
); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -17138,6 +17144,7 @@ INSERT INTO mcp_server_configs (
|
||||
oauth2_client_secret_key_id,
|
||||
oauth2_auth_url,
|
||||
oauth2_token_url,
|
||||
oauth2_revocation_url,
|
||||
oauth2_scopes,
|
||||
api_key_header,
|
||||
api_key_value,
|
||||
@@ -17172,18 +17179,19 @@ INSERT INTO mcp_server_configs (
|
||||
$16::text,
|
||||
$17::text,
|
||||
$18::text,
|
||||
$19::text[],
|
||||
$19::text,
|
||||
$20::text[],
|
||||
$21::text,
|
||||
$22::boolean,
|
||||
$21::text[],
|
||||
$22::text,
|
||||
$23::boolean,
|
||||
$24::boolean,
|
||||
$25::boolean,
|
||||
$26::uuid,
|
||||
$27::uuid
|
||||
$26::boolean,
|
||||
$27::uuid,
|
||||
$28::uuid
|
||||
)
|
||||
RETURNING
|
||||
id, display_name, slug, description, icon_url, transport, url, auth_type, oauth2_client_id, oauth2_client_secret, oauth2_client_secret_key_id, oauth2_auth_url, oauth2_token_url, oauth2_scopes, api_key_header, api_key_value, api_key_value_key_id, custom_headers, custom_headers_key_id, tool_allow_list, tool_deny_list, availability, enabled, created_by, updated_by, created_at, updated_at, model_intent, allow_in_plan_mode, forward_coder_headers
|
||||
id, display_name, slug, description, icon_url, transport, url, auth_type, oauth2_client_id, oauth2_client_secret, oauth2_client_secret_key_id, oauth2_auth_url, oauth2_token_url, oauth2_scopes, api_key_header, api_key_value, api_key_value_key_id, custom_headers, custom_headers_key_id, tool_allow_list, tool_deny_list, availability, enabled, created_by, updated_by, created_at, updated_at, model_intent, allow_in_plan_mode, forward_coder_headers, oauth2_revocation_url
|
||||
`
|
||||
|
||||
type InsertMCPServerConfigParams struct {
|
||||
@@ -17199,6 +17207,7 @@ type InsertMCPServerConfigParams struct {
|
||||
OAuth2ClientSecretKeyID sql.NullString `db:"oauth2_client_secret_key_id" json:"oauth2_client_secret_key_id"`
|
||||
OAuth2AuthURL string `db:"oauth2_auth_url" json:"oauth2_auth_url"`
|
||||
OAuth2TokenURL string `db:"oauth2_token_url" json:"oauth2_token_url"`
|
||||
OAuth2RevocationURL string `db:"oauth2_revocation_url" json:"oauth2_revocation_url"`
|
||||
OAuth2Scopes string `db:"oauth2_scopes" json:"oauth2_scopes"`
|
||||
APIKeyHeader string `db:"api_key_header" json:"api_key_header"`
|
||||
APIKeyValue string `db:"api_key_value" json:"api_key_value"`
|
||||
@@ -17230,6 +17239,7 @@ func (q *sqlQuerier) InsertMCPServerConfig(ctx context.Context, arg InsertMCPSer
|
||||
arg.OAuth2ClientSecretKeyID,
|
||||
arg.OAuth2AuthURL,
|
||||
arg.OAuth2TokenURL,
|
||||
arg.OAuth2RevocationURL,
|
||||
arg.OAuth2Scopes,
|
||||
arg.APIKeyHeader,
|
||||
arg.APIKeyValue,
|
||||
@@ -17278,6 +17288,7 @@ func (q *sqlQuerier) InsertMCPServerConfig(ctx context.Context, arg InsertMCPSer
|
||||
&i.ModelIntent,
|
||||
&i.AllowInPlanMode,
|
||||
&i.ForwardCoderHeaders,
|
||||
&i.OAuth2RevocationURL,
|
||||
)
|
||||
return i, err
|
||||
}
|
||||
@@ -17346,25 +17357,26 @@ SET
|
||||
oauth2_client_secret_key_id = $10::text,
|
||||
oauth2_auth_url = $11::text,
|
||||
oauth2_token_url = $12::text,
|
||||
oauth2_scopes = $13::text,
|
||||
api_key_header = $14::text,
|
||||
api_key_value = $15::text,
|
||||
api_key_value_key_id = $16::text,
|
||||
custom_headers = $17::text,
|
||||
custom_headers_key_id = $18::text,
|
||||
tool_allow_list = $19::text[],
|
||||
tool_deny_list = $20::text[],
|
||||
availability = $21::text,
|
||||
enabled = $22::boolean,
|
||||
model_intent = $23::boolean,
|
||||
allow_in_plan_mode = $24::boolean,
|
||||
forward_coder_headers = $25::boolean,
|
||||
updated_by = $26::uuid,
|
||||
oauth2_revocation_url = $13::text,
|
||||
oauth2_scopes = $14::text,
|
||||
api_key_header = $15::text,
|
||||
api_key_value = $16::text,
|
||||
api_key_value_key_id = $17::text,
|
||||
custom_headers = $18::text,
|
||||
custom_headers_key_id = $19::text,
|
||||
tool_allow_list = $20::text[],
|
||||
tool_deny_list = $21::text[],
|
||||
availability = $22::text,
|
||||
enabled = $23::boolean,
|
||||
model_intent = $24::boolean,
|
||||
allow_in_plan_mode = $25::boolean,
|
||||
forward_coder_headers = $26::boolean,
|
||||
updated_by = $27::uuid,
|
||||
updated_at = NOW()
|
||||
WHERE
|
||||
id = $27::uuid
|
||||
id = $28::uuid
|
||||
RETURNING
|
||||
id, display_name, slug, description, icon_url, transport, url, auth_type, oauth2_client_id, oauth2_client_secret, oauth2_client_secret_key_id, oauth2_auth_url, oauth2_token_url, oauth2_scopes, api_key_header, api_key_value, api_key_value_key_id, custom_headers, custom_headers_key_id, tool_allow_list, tool_deny_list, availability, enabled, created_by, updated_by, created_at, updated_at, model_intent, allow_in_plan_mode, forward_coder_headers
|
||||
id, display_name, slug, description, icon_url, transport, url, auth_type, oauth2_client_id, oauth2_client_secret, oauth2_client_secret_key_id, oauth2_auth_url, oauth2_token_url, oauth2_scopes, api_key_header, api_key_value, api_key_value_key_id, custom_headers, custom_headers_key_id, tool_allow_list, tool_deny_list, availability, enabled, created_by, updated_by, created_at, updated_at, model_intent, allow_in_plan_mode, forward_coder_headers, oauth2_revocation_url
|
||||
`
|
||||
|
||||
type UpdateMCPServerConfigParams struct {
|
||||
@@ -17380,6 +17392,7 @@ type UpdateMCPServerConfigParams struct {
|
||||
OAuth2ClientSecretKeyID sql.NullString `db:"oauth2_client_secret_key_id" json:"oauth2_client_secret_key_id"`
|
||||
OAuth2AuthURL string `db:"oauth2_auth_url" json:"oauth2_auth_url"`
|
||||
OAuth2TokenURL string `db:"oauth2_token_url" json:"oauth2_token_url"`
|
||||
OAuth2RevocationURL string `db:"oauth2_revocation_url" json:"oauth2_revocation_url"`
|
||||
OAuth2Scopes string `db:"oauth2_scopes" json:"oauth2_scopes"`
|
||||
APIKeyHeader string `db:"api_key_header" json:"api_key_header"`
|
||||
APIKeyValue string `db:"api_key_value" json:"api_key_value"`
|
||||
@@ -17411,6 +17424,7 @@ func (q *sqlQuerier) UpdateMCPServerConfig(ctx context.Context, arg UpdateMCPSer
|
||||
arg.OAuth2ClientSecretKeyID,
|
||||
arg.OAuth2AuthURL,
|
||||
arg.OAuth2TokenURL,
|
||||
arg.OAuth2RevocationURL,
|
||||
arg.OAuth2Scopes,
|
||||
arg.APIKeyHeader,
|
||||
arg.APIKeyValue,
|
||||
@@ -17459,6 +17473,67 @@ func (q *sqlQuerier) UpdateMCPServerConfig(ctx context.Context, arg UpdateMCPSer
|
||||
&i.ModelIntent,
|
||||
&i.AllowInPlanMode,
|
||||
&i.ForwardCoderHeaders,
|
||||
&i.OAuth2RevocationURL,
|
||||
)
|
||||
return i, err
|
||||
}
|
||||
|
||||
const updateMCPServerUserTokenFromRefresh = `-- name: UpdateMCPServerUserTokenFromRefresh :one
|
||||
UPDATE mcp_server_user_tokens
|
||||
SET
|
||||
access_token = $1::text,
|
||||
access_token_key_id = $2::text,
|
||||
refresh_token = $3::text,
|
||||
refresh_token_key_id = $4::text,
|
||||
token_type = $5::text,
|
||||
expiry = $6::timestamptz,
|
||||
oauth_refresh_failure_reason = '',
|
||||
updated_at = NOW()
|
||||
WHERE
|
||||
id = $7::uuid
|
||||
AND updated_at = $8::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 UpdateMCPServerUserTokenFromRefreshParams struct {
|
||||
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"`
|
||||
ID uuid.UUID `db:"id" json:"id"`
|
||||
UpdatedAt time.Time `db:"updated_at" json:"updated_at"`
|
||||
}
|
||||
|
||||
// Refresh persistence must not recreate a token deleted by disconnect.
|
||||
// The optimistic lock also prevents stale refreshes from replacing newer tokens.
|
||||
func (q *sqlQuerier) UpdateMCPServerUserTokenFromRefresh(ctx context.Context, arg UpdateMCPServerUserTokenFromRefreshParams) (MCPServerUserToken, error) {
|
||||
row := q.db.QueryRowContext(ctx, updateMCPServerUserTokenFromRefresh,
|
||||
arg.AccessToken,
|
||||
arg.AccessTokenKeyID,
|
||||
arg.RefreshToken,
|
||||
arg.RefreshTokenKeyID,
|
||||
arg.TokenType,
|
||||
arg.Expiry,
|
||||
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
|
||||
}
|
||||
|
||||
@@ -67,6 +67,7 @@ INSERT INTO mcp_server_configs (
|
||||
oauth2_client_secret_key_id,
|
||||
oauth2_auth_url,
|
||||
oauth2_token_url,
|
||||
oauth2_revocation_url,
|
||||
oauth2_scopes,
|
||||
api_key_header,
|
||||
api_key_value,
|
||||
@@ -95,6 +96,7 @@ INSERT INTO mcp_server_configs (
|
||||
sqlc.narg('oauth2_client_secret_key_id')::text,
|
||||
@oauth2_auth_url::text,
|
||||
@oauth2_token_url::text,
|
||||
@oauth2_revocation_url::text,
|
||||
@oauth2_scopes::text,
|
||||
@api_key_header::text,
|
||||
@api_key_value::text,
|
||||
@@ -130,6 +132,7 @@ SET
|
||||
oauth2_client_secret_key_id = sqlc.narg('oauth2_client_secret_key_id')::text,
|
||||
oauth2_auth_url = @oauth2_auth_url::text,
|
||||
oauth2_token_url = @oauth2_token_url::text,
|
||||
oauth2_revocation_url = @oauth2_revocation_url::text,
|
||||
oauth2_scopes = @oauth2_scopes::text,
|
||||
api_key_header = @api_key_header::text,
|
||||
api_key_value = @api_key_value::text,
|
||||
@@ -207,6 +210,25 @@ ON CONFLICT (mcp_server_config_id, user_id) DO UPDATE SET
|
||||
RETURNING
|
||||
*;
|
||||
|
||||
-- name: UpdateMCPServerUserTokenFromRefresh :one
|
||||
-- Refresh persistence must not recreate a token deleted by disconnect.
|
||||
-- The optimistic lock also prevents stale refreshes from replacing newer tokens.
|
||||
UPDATE mcp_server_user_tokens
|
||||
SET
|
||||
access_token = @access_token::text,
|
||||
access_token_key_id = sqlc.narg('access_token_key_id')::text,
|
||||
refresh_token = @refresh_token::text,
|
||||
refresh_token_key_id = sqlc.narg('refresh_token_key_id')::text,
|
||||
token_type = @token_type::text,
|
||||
expiry = sqlc.narg('expiry')::timestamptz,
|
||||
oauth_refresh_failure_reason = '',
|
||||
updated_at = NOW()
|
||||
WHERE
|
||||
id = @id::uuid
|
||||
AND updated_at = @updated_at::timestamptz
|
||||
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.
|
||||
|
||||
@@ -288,6 +288,7 @@ sql:
|
||||
oauth2_client_secret_key_id: OAuth2ClientSecretKeyID
|
||||
oauth2_auth_url: OAuth2AuthURL
|
||||
oauth2_token_url: OAuth2TokenURL
|
||||
oauth2_revocation_url: OAuth2RevocationURL
|
||||
oauth2_scopes: OAuth2Scopes
|
||||
api_key_header: APIKeyHeader
|
||||
api_key_value: APIKeyValue
|
||||
|
||||
+176
-42
@@ -195,7 +195,7 @@ func (api *API) listMCPServerConfigs(rw http.ResponseWriter, r *http.Request) {
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
tokenMap[tok.MCPServerConfigID] = api.refreshMCPUserToken(ctx, cfg, tok, apiKey.UserID)
|
||||
tokenMap[tok.MCPServerConfigID] = api.refreshMCPUserToken(ctx, cfg, tok)
|
||||
}
|
||||
|
||||
resp := make([]codersdk.MCPServerConfig, 0, len(configs))
|
||||
@@ -233,6 +233,16 @@ func (api *API) createMCPServerConfig(rw http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
if trimmed := strings.TrimSpace(req.OAuth2RevocationURL); trimmed != "" {
|
||||
if err := mcpclient.ValidateRevocationEndpoint(trimmed); err != nil {
|
||||
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
|
||||
Message: "Invalid OAuth2 revocation URL.",
|
||||
Detail: "oauth2_revocation_url must be an https URL (loopback hosts may use http).",
|
||||
})
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// Validate auth-type-dependent fields.
|
||||
switch req.AuthType {
|
||||
case "oauth2":
|
||||
@@ -269,6 +279,7 @@ func (api *API) createMCPServerConfig(rw http.ResponseWriter, r *http.Request) {
|
||||
OAuth2ClientSecretKeyID: sql.NullString{},
|
||||
OAuth2AuthURL: "",
|
||||
OAuth2TokenURL: "",
|
||||
OAuth2RevocationURL: "",
|
||||
OAuth2Scopes: "",
|
||||
APIKeyHeader: strings.TrimSpace(req.APIKeyHeader),
|
||||
APIKeyValue: strings.TrimSpace(req.APIKeyValue),
|
||||
@@ -343,6 +354,22 @@ func (api *API) createMCPServerConfig(rw http.ResponseWriter, r *http.Request) {
|
||||
oauth2Scopes = result.scopes
|
||||
}
|
||||
|
||||
// A discovered endpoint that fails the HTTPS policy is
|
||||
// dropped instead of failing creation.
|
||||
oauth2RevocationURL := strings.TrimSpace(req.OAuth2RevocationURL)
|
||||
if oauth2RevocationURL == "" {
|
||||
oauth2RevocationURL = result.revocationURL
|
||||
if oauth2RevocationURL != "" {
|
||||
if err := mcpclient.ValidateRevocationEndpoint(oauth2RevocationURL); err != nil {
|
||||
api.Logger.Warn(ctx, "ignoring discovered MCP oauth2 revocation endpoint",
|
||||
slog.F("url", req.URL),
|
||||
slog.Error(err),
|
||||
)
|
||||
oauth2RevocationURL = ""
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Update the record with discovered OAuth2 credentials.
|
||||
updated, err := api.Database.UpdateMCPServerConfig(ctx, database.UpdateMCPServerConfigParams{
|
||||
ID: inserted.ID,
|
||||
@@ -358,6 +385,7 @@ func (api *API) createMCPServerConfig(rw http.ResponseWriter, r *http.Request) {
|
||||
OAuth2ClientSecretKeyID: sql.NullString{},
|
||||
OAuth2AuthURL: result.authURL,
|
||||
OAuth2TokenURL: result.tokenURL,
|
||||
OAuth2RevocationURL: oauth2RevocationURL,
|
||||
OAuth2Scopes: oauth2Scopes,
|
||||
APIKeyHeader: inserted.APIKeyHeader,
|
||||
APIKeyValue: inserted.APIKeyValue,
|
||||
@@ -428,6 +456,7 @@ func (api *API) createMCPServerConfig(rw http.ResponseWriter, r *http.Request) {
|
||||
OAuth2ClientSecretKeyID: sql.NullString{},
|
||||
OAuth2AuthURL: strings.TrimSpace(req.OAuth2AuthURL),
|
||||
OAuth2TokenURL: strings.TrimSpace(req.OAuth2TokenURL),
|
||||
OAuth2RevocationURL: strings.TrimSpace(req.OAuth2RevocationURL),
|
||||
OAuth2Scopes: strings.TrimSpace(req.OAuth2Scopes),
|
||||
APIKeyHeader: strings.TrimSpace(req.APIKeyHeader),
|
||||
APIKeyValue: strings.TrimSpace(req.APIKeyValue),
|
||||
@@ -531,7 +560,7 @@ func (api *API) getMCPServerConfig(rw http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
for _, tok := range userTokens {
|
||||
if tok.MCPServerConfigID == config.ID {
|
||||
sdkConfig.AuthConnected = api.refreshMCPUserToken(ctx, config, tok, apiKey.UserID)
|
||||
sdkConfig.AuthConnected = api.refreshMCPUserToken(ctx, config, tok)
|
||||
break
|
||||
}
|
||||
}
|
||||
@@ -563,6 +592,29 @@ func (api *API) updateMCPServerConfig(rw http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
// Validated here rather than via a struct tag because an empty
|
||||
// string is a valid value that clears the stored URL.
|
||||
if req.OAuth2RevocationURL != nil {
|
||||
if trimmed := strings.TrimSpace(*req.OAuth2RevocationURL); trimmed != "" {
|
||||
if err := httpapi.Validate.VarCtx(ctx, trimmed, "url"); err != nil {
|
||||
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
|
||||
Message: "Invalid OAuth2 revocation URL.",
|
||||
Detail: "oauth2_revocation_url must be a valid URL or an empty string.",
|
||||
})
|
||||
return
|
||||
}
|
||||
// Same policy as RevokeOAuth2Token, so stored URLs are
|
||||
// not refused later at disconnect time.
|
||||
if err := mcpclient.ValidateRevocationEndpoint(trimmed); err != nil {
|
||||
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
|
||||
Message: "Invalid OAuth2 revocation URL.",
|
||||
Detail: "oauth2_revocation_url must be an https URL (loopback hosts may use http).",
|
||||
})
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Pre-validate custom headers before entering the transaction.
|
||||
var customHeadersJSON string
|
||||
if req.CustomHeaders != nil {
|
||||
@@ -642,6 +694,11 @@ func (api *API) updateMCPServerConfig(rw http.ResponseWriter, r *http.Request) {
|
||||
oauth2TokenURL = strings.TrimSpace(*req.OAuth2TokenURL)
|
||||
}
|
||||
|
||||
oauth2RevocationURL := existing.OAuth2RevocationURL
|
||||
if req.OAuth2RevocationURL != nil {
|
||||
oauth2RevocationURL = strings.TrimSpace(*req.OAuth2RevocationURL)
|
||||
}
|
||||
|
||||
oauth2Scopes := existing.OAuth2Scopes
|
||||
if req.OAuth2Scopes != nil {
|
||||
oauth2Scopes = strings.TrimSpace(*req.OAuth2Scopes)
|
||||
@@ -713,6 +770,7 @@ func (api *API) updateMCPServerConfig(rw http.ResponseWriter, r *http.Request) {
|
||||
oauth2ClientSecretKeyID = sql.NullString{}
|
||||
oauth2AuthURL = ""
|
||||
oauth2TokenURL = ""
|
||||
oauth2RevocationURL = ""
|
||||
oauth2Scopes = ""
|
||||
apiKeyHeader = ""
|
||||
apiKeyValue = ""
|
||||
@@ -731,6 +789,7 @@ func (api *API) updateMCPServerConfig(rw http.ResponseWriter, r *http.Request) {
|
||||
oauth2ClientSecretKeyID = sql.NullString{}
|
||||
oauth2AuthURL = ""
|
||||
oauth2TokenURL = ""
|
||||
oauth2RevocationURL = ""
|
||||
oauth2Scopes = ""
|
||||
customHeaders = "{}"
|
||||
customHeadersKeyID = sql.NullString{}
|
||||
@@ -740,6 +799,7 @@ func (api *API) updateMCPServerConfig(rw http.ResponseWriter, r *http.Request) {
|
||||
oauth2ClientSecretKeyID = sql.NullString{}
|
||||
oauth2AuthURL = ""
|
||||
oauth2TokenURL = ""
|
||||
oauth2RevocationURL = ""
|
||||
oauth2Scopes = ""
|
||||
apiKeyHeader = ""
|
||||
apiKeyValue = ""
|
||||
@@ -753,6 +813,7 @@ func (api *API) updateMCPServerConfig(rw http.ResponseWriter, r *http.Request) {
|
||||
oauth2ClientSecretKeyID = sql.NullString{}
|
||||
oauth2AuthURL = ""
|
||||
oauth2TokenURL = ""
|
||||
oauth2RevocationURL = ""
|
||||
oauth2Scopes = ""
|
||||
apiKeyHeader = ""
|
||||
apiKeyValue = ""
|
||||
@@ -775,6 +836,7 @@ func (api *API) updateMCPServerConfig(rw http.ResponseWriter, r *http.Request) {
|
||||
OAuth2ClientSecretKeyID: oauth2ClientSecretKeyID,
|
||||
OAuth2AuthURL: oauth2AuthURL,
|
||||
OAuth2TokenURL: oauth2TokenURL,
|
||||
OAuth2RevocationURL: oauth2RevocationURL,
|
||||
OAuth2Scopes: oauth2Scopes,
|
||||
APIKeyHeader: apiKeyHeader,
|
||||
APIKeyValue: apiKeyValue,
|
||||
@@ -1138,6 +1200,7 @@ func (api *API) mcpServerOAuth2Callback(rw http.ResponseWriter, r *http.Request)
|
||||
// @x-apidocgen {"skip": true}
|
||||
// EXPERIMENTAL: this endpoint is experimental and is subject to change.
|
||||
// Removes the user's stored OAuth2 token for an MCP server.
|
||||
// Provider revocation is best-effort and cannot block local deletion.
|
||||
func (api *API) mcpServerOAuth2Disconnect(rw http.ResponseWriter, r *http.Request) {
|
||||
ctx := r.Context()
|
||||
apiKey := httpmw.APIKey(r)
|
||||
@@ -1148,11 +1211,43 @@ func (api *API) mcpServerOAuth2Disconnect(rw http.ResponseWriter, r *http.Reques
|
||||
}
|
||||
|
||||
//nolint:gocritic // Users manage their own tokens.
|
||||
err := api.Database.DeleteMCPServerUserToken(dbauthz.AsSystemRestricted(ctx), database.DeleteMCPServerUserTokenParams{
|
||||
MCPServerConfigID: mcpServerID,
|
||||
UserID: apiKey.UserID,
|
||||
})
|
||||
systemCtx := dbauthz.AsSystemRestricted(ctx)
|
||||
var (
|
||||
config database.MCPServerConfig
|
||||
token database.MCPServerUserToken
|
||||
)
|
||||
// Serializable isolation keeps the revoked token aligned with the row deleted locally.
|
||||
err := api.Database.InTx(func(tx database.Store) error {
|
||||
dbToken, err := tx.GetMCPServerUserToken(systemCtx, database.GetMCPServerUserTokenParams{
|
||||
MCPServerConfigID: mcpServerID,
|
||||
UserID: apiKey.UserID,
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
// Load the config only after the token is found so callers
|
||||
// without a token cannot probe which config IDs exist.
|
||||
dbConfig, err := tx.GetMCPServerConfigByID(systemCtx, mcpServerID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.DeleteMCPServerUserToken(systemCtx, database.DeleteMCPServerUserTokenParams{
|
||||
MCPServerConfigID: mcpServerID,
|
||||
UserID: apiKey.UserID,
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
config = dbConfig
|
||||
token = dbToken
|
||||
return nil
|
||||
}, &database.TxOptions{Isolation: sql.LevelSerializable})
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
// Nonexistent config IDs take the same path, so they
|
||||
// cannot be probed either.
|
||||
httpapi.Write(ctx, rw, http.StatusOK, codersdk.MCPServerOAuth2DisconnectResponse{})
|
||||
return
|
||||
}
|
||||
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
|
||||
Message: "Failed to disconnect OAuth2 token.",
|
||||
Detail: err.Error(),
|
||||
@@ -1160,7 +1255,24 @@ func (api *API) mcpServerOAuth2Disconnect(rw http.ResponseWriter, r *http.Reques
|
||||
return
|
||||
}
|
||||
|
||||
rw.WriteHeader(http.StatusNoContent)
|
||||
resp := codersdk.MCPServerOAuth2DisconnectResponse{}
|
||||
if config.AuthType == "oauth2" {
|
||||
// The local token is already deleted, so a client abort must
|
||||
// not cancel the provider revocation; it has its own timeout.
|
||||
revoked, err := mcpclient.RevokeOAuth2Token(context.WithoutCancel(ctx), api.HTTPClient, config, token)
|
||||
resp.TokenRevoked = revoked
|
||||
if err != nil {
|
||||
api.Logger.Warn(ctx, "failed to revoke MCP oauth2 token at provider",
|
||||
slog.F("server_slug", config.Slug),
|
||||
slog.Error(err),
|
||||
)
|
||||
// Provider error bodies may echo the client secret, so
|
||||
// callers only get a generic message.
|
||||
resp.TokenRevocationError = "The OAuth provider rejected the revocation request."
|
||||
}
|
||||
}
|
||||
|
||||
httpapi.Write(ctx, rw, http.StatusOK, resp)
|
||||
}
|
||||
|
||||
// refreshMCPUserToken attempts to refresh an expired OAuth2 token
|
||||
@@ -1173,7 +1285,6 @@ func (api *API) refreshMCPUserToken(
|
||||
ctx context.Context,
|
||||
cfg database.MCPServerConfig,
|
||||
tok database.MCPServerUserToken,
|
||||
userID uuid.UUID,
|
||||
) bool {
|
||||
if cfg.AuthType != "oauth2" {
|
||||
return true
|
||||
@@ -1209,11 +1320,11 @@ func (api *API) refreshMCPUserToken(
|
||||
|
||||
//nolint:gocritic // Need system-level write access to
|
||||
// persist the refreshed OAuth2 token.
|
||||
_, err = api.Database.UpsertMCPServerUserToken(
|
||||
_, err = api.Database.UpdateMCPServerUserTokenFromRefresh(
|
||||
dbauthz.AsSystemRestricted(ctx),
|
||||
database.UpsertMCPServerUserTokenParams{
|
||||
MCPServerConfigID: tok.MCPServerConfigID,
|
||||
UserID: userID,
|
||||
database.UpdateMCPServerUserTokenFromRefreshParams{
|
||||
ID: tok.ID,
|
||||
UpdatedAt: tok.UpdatedAt,
|
||||
AccessToken: result.AccessToken,
|
||||
AccessTokenKeyID: sql.NullString{},
|
||||
RefreshToken: result.RefreshToken,
|
||||
@@ -1223,6 +1334,13 @@ func (api *API) refreshMCPUserToken(
|
||||
},
|
||||
)
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
connected, readErr := api.currentMCPUserTokenConnected(ctx, tok)
|
||||
if readErr == nil {
|
||||
return connected
|
||||
}
|
||||
err = readErr
|
||||
}
|
||||
api.Logger.Warn(ctx, "failed to persist refreshed MCP oauth2 token",
|
||||
slog.F("server_slug", cfg.Slug),
|
||||
slog.Error(err),
|
||||
@@ -1233,6 +1351,29 @@ func (api *API) refreshMCPUserToken(
|
||||
return true
|
||||
}
|
||||
|
||||
func (api *API) currentMCPUserTokenConnected(
|
||||
ctx context.Context,
|
||||
tok database.MCPServerUserToken,
|
||||
) (bool, error) {
|
||||
//nolint:gocritic // Reading the current token requires system access.
|
||||
current, err := api.Database.GetMCPServerUserToken(
|
||||
dbauthz.AsSystemRestricted(ctx),
|
||||
database.GetMCPServerUserTokenParams{
|
||||
MCPServerConfigID: tok.MCPServerConfigID,
|
||||
UserID: tok.UserID,
|
||||
},
|
||||
)
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return false, nil
|
||||
}
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
return current.OauthRefreshFailureReason == "" &&
|
||||
current.AccessToken != "" &&
|
||||
(!current.Expiry.Valid || current.Expiry.Time.After(time.Now())), nil
|
||||
}
|
||||
|
||||
// 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
|
||||
@@ -1258,21 +1399,9 @@ func (api *API) markMCPTokenRefreshFailure(
|
||||
}
|
||||
|
||||
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,
|
||||
},
|
||||
)
|
||||
connected, readErr := api.currentMCPUserTokenConnected(ctx, tok)
|
||||
if readErr == nil {
|
||||
return current.OauthRefreshFailureReason == "" &&
|
||||
current.AccessToken != "" &&
|
||||
(!current.Expiry.Valid || current.Expiry.Time.After(time.Now()))
|
||||
return connected
|
||||
}
|
||||
err = readErr
|
||||
}
|
||||
@@ -1312,12 +1441,13 @@ func convertMCPServerConfig(config database.MCPServerConfig) codersdk.MCPServerC
|
||||
Transport: config.Transport,
|
||||
URL: config.Url,
|
||||
|
||||
AuthType: config.AuthType,
|
||||
OAuth2ClientID: config.OAuth2ClientID,
|
||||
HasOAuth2Secret: config.OAuth2ClientSecret != "",
|
||||
OAuth2AuthURL: config.OAuth2AuthURL,
|
||||
OAuth2TokenURL: config.OAuth2TokenURL,
|
||||
OAuth2Scopes: config.OAuth2Scopes,
|
||||
AuthType: config.AuthType,
|
||||
OAuth2ClientID: config.OAuth2ClientID,
|
||||
HasOAuth2Secret: config.OAuth2ClientSecret != "",
|
||||
OAuth2AuthURL: config.OAuth2AuthURL,
|
||||
OAuth2TokenURL: config.OAuth2TokenURL,
|
||||
OAuth2RevocationURL: config.OAuth2RevocationURL,
|
||||
OAuth2Scopes: config.OAuth2Scopes,
|
||||
|
||||
APIKeyHeader: config.APIKeyHeader,
|
||||
HasAPIKey: config.APIKeyValue != "",
|
||||
@@ -1352,6 +1482,7 @@ func convertMCPServerConfigRedacted(config database.MCPServerConfig) codersdk.MC
|
||||
c.OAuth2ClientID = ""
|
||||
c.OAuth2AuthURL = ""
|
||||
c.OAuth2TokenURL = ""
|
||||
c.OAuth2RevocationURL = ""
|
||||
c.OAuth2Scopes = ""
|
||||
c.APIKeyHeader = ""
|
||||
return c
|
||||
@@ -1398,11 +1529,12 @@ func coalesceStringSlice(ss []string) []string {
|
||||
// mcpOAuth2Discovery holds the result of MCP OAuth2 auto-discovery
|
||||
// and Dynamic Client Registration.
|
||||
type mcpOAuth2Discovery struct {
|
||||
clientID string
|
||||
clientSecret string
|
||||
authURL string
|
||||
tokenURL string
|
||||
scopes string // space-separated
|
||||
clientID string
|
||||
clientSecret string
|
||||
authURL string
|
||||
tokenURL string
|
||||
revocationURL string
|
||||
scopes string // space-separated
|
||||
}
|
||||
|
||||
// protectedResourceMetadata represents the response from a
|
||||
@@ -1420,6 +1552,7 @@ type authServerMetadata struct {
|
||||
AuthorizationEndpoint string `json:"authorization_endpoint"`
|
||||
TokenEndpoint string `json:"token_endpoint"`
|
||||
RegistrationEndpoint string `json:"registration_endpoint,omitempty"`
|
||||
RevocationEndpoint string `json:"revocation_endpoint,omitempty"`
|
||||
ScopesSupported []string `json:"scopes_supported,omitempty"`
|
||||
}
|
||||
|
||||
@@ -1740,10 +1873,11 @@ func discoverAndRegisterMCPOAuth2(ctx context.Context, httpClient *http.Client,
|
||||
scopes := strings.Join(asMeta.ScopesSupported, " ")
|
||||
|
||||
return &mcpOAuth2Discovery{
|
||||
clientID: clientID,
|
||||
clientSecret: clientSecret,
|
||||
authURL: asMeta.AuthorizationEndpoint,
|
||||
tokenURL: asMeta.TokenEndpoint,
|
||||
scopes: scopes,
|
||||
clientID: clientID,
|
||||
clientSecret: clientSecret,
|
||||
authURL: asMeta.AuthorizationEndpoint,
|
||||
tokenURL: asMeta.TokenEndpoint,
|
||||
revocationURL: asMeta.RevocationEndpoint,
|
||||
scopes: scopes,
|
||||
}, nil
|
||||
}
|
||||
|
||||
+422
-76
@@ -7,7 +7,9 @@ import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/url"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
@@ -208,23 +210,24 @@ func TestMCPServerConfigsSecretsNeverLeaked(t *testing.T) {
|
||||
|
||||
// Create a config with ALL secret fields populated.
|
||||
created, err := adminClient.CreateMCPServerConfig(ctx, codersdk.CreateMCPServerConfigRequest{
|
||||
DisplayName: "Secrets Test",
|
||||
Slug: "secrets-test",
|
||||
Transport: "streamable_http",
|
||||
URL: "https://mcp.example.com/secrets",
|
||||
AuthType: "oauth2",
|
||||
OAuth2ClientID: "client-id-secret-test",
|
||||
OAuth2ClientSecret: "THIS-IS-A-SECRET-VALUE",
|
||||
OAuth2AuthURL: "https://auth.example.com/authorize",
|
||||
OAuth2TokenURL: "https://auth.example.com/token",
|
||||
OAuth2Scopes: "read write",
|
||||
APIKeyHeader: "X-Api-Key",
|
||||
APIKeyValue: "THIS-IS-A-SECRET-API-KEY",
|
||||
CustomHeaders: map[string]string{"X-Custom": "THIS-IS-A-SECRET-HEADER"},
|
||||
Availability: "default_on",
|
||||
Enabled: true,
|
||||
ToolAllowList: []string{},
|
||||
ToolDenyList: []string{},
|
||||
DisplayName: "Secrets Test",
|
||||
Slug: "secrets-test",
|
||||
Transport: "streamable_http",
|
||||
URL: "https://mcp.example.com/secrets",
|
||||
AuthType: "oauth2",
|
||||
OAuth2ClientID: "client-id-secret-test",
|
||||
OAuth2ClientSecret: "THIS-IS-A-SECRET-VALUE",
|
||||
OAuth2AuthURL: "https://auth.example.com/authorize",
|
||||
OAuth2TokenURL: "https://auth.example.com/token",
|
||||
OAuth2RevocationURL: "https://auth.example.com/revoke",
|
||||
OAuth2Scopes: "read write",
|
||||
APIKeyHeader: "X-Api-Key",
|
||||
APIKeyValue: "THIS-IS-A-SECRET-API-KEY",
|
||||
CustomHeaders: map[string]string{"X-Custom": "THIS-IS-A-SECRET-HEADER"},
|
||||
Availability: "default_on",
|
||||
Enabled: true,
|
||||
ToolAllowList: []string{},
|
||||
ToolDenyList: []string{},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
@@ -277,6 +280,7 @@ func TestMCPServerConfigsSecretsNeverLeaked(t *testing.T) {
|
||||
assert.Empty(t, cfg.OAuth2ClientID, "member should not see OAuth2ClientID")
|
||||
assert.Empty(t, cfg.OAuth2AuthURL, "member should not see OAuth2AuthURL")
|
||||
assert.Empty(t, cfg.OAuth2TokenURL, "member should not see OAuth2TokenURL")
|
||||
assert.Empty(t, cfg.OAuth2RevocationURL, "member should not see OAuth2RevocationURL")
|
||||
assert.Empty(t, cfg.APIKeyHeader, "member should not see APIKeyHeader")
|
||||
assert.Empty(t, cfg.OAuth2Scopes, "member should not see OAuth2Scopes")
|
||||
assert.Empty(t, cfg.URL, "member should not see URL")
|
||||
@@ -373,27 +377,85 @@ func TestMCPServerConfigsUserOIDCClearsFields(t *testing.T) {
|
||||
// switch the auth_type to user_oidc and verify all auth-specific
|
||||
// fields are cleared.
|
||||
created, err := client.CreateMCPServerConfig(ctx, codersdk.CreateMCPServerConfigRequest{
|
||||
DisplayName: "Switch Server",
|
||||
Slug: "switch-server",
|
||||
Transport: "streamable_http",
|
||||
URL: "https://mcp.example.com/v1",
|
||||
AuthType: "oauth2",
|
||||
OAuth2ClientID: "cid",
|
||||
OAuth2ClientSecret: "secret-value",
|
||||
OAuth2AuthURL: "https://auth.example.com/authorize",
|
||||
OAuth2TokenURL: "https://auth.example.com/token",
|
||||
OAuth2Scopes: "read write",
|
||||
Availability: "default_off",
|
||||
Enabled: true,
|
||||
ToolAllowList: []string{},
|
||||
ToolDenyList: []string{},
|
||||
DisplayName: "Switch Server",
|
||||
Slug: "switch-server",
|
||||
Transport: "streamable_http",
|
||||
URL: "https://mcp.example.com/v1",
|
||||
AuthType: "oauth2",
|
||||
OAuth2ClientID: "cid",
|
||||
OAuth2ClientSecret: "secret-value",
|
||||
OAuth2AuthURL: "https://auth.example.com/authorize",
|
||||
OAuth2TokenURL: "https://auth.example.com/token",
|
||||
OAuth2RevocationURL: "https://auth.example.com/revoke",
|
||||
OAuth2Scopes: "read write",
|
||||
Availability: "default_off",
|
||||
Enabled: true,
|
||||
ToolAllowList: []string{},
|
||||
ToolDenyList: []string{},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.True(t, created.HasOAuth2Secret)
|
||||
require.Equal(t, "cid", created.OAuth2ClientID)
|
||||
require.Equal(t, "https://auth.example.com/revoke", created.OAuth2RevocationURL)
|
||||
|
||||
newRevocationURL := "https://auth.example.com/revoke2"
|
||||
updated, err := client.UpdateMCPServerConfig(ctx, created.ID, codersdk.UpdateMCPServerConfigRequest{
|
||||
OAuth2RevocationURL: &newRevocationURL,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, newRevocationURL, updated.OAuth2RevocationURL)
|
||||
|
||||
invalidURL := "not a url"
|
||||
_, err = client.UpdateMCPServerConfig(ctx, created.ID, codersdk.UpdateMCPServerConfigRequest{
|
||||
OAuth2RevocationURL: &invalidURL,
|
||||
})
|
||||
require.Error(t, err)
|
||||
var sdkErr *codersdk.Error
|
||||
require.ErrorAs(t, err, &sdkErr)
|
||||
require.Equal(t, http.StatusBadRequest, sdkErr.StatusCode())
|
||||
|
||||
// Plaintext URLs are rejected on save, not later at disconnect.
|
||||
plaintextURL := "http://auth.example.com/revoke"
|
||||
_, err = client.UpdateMCPServerConfig(ctx, created.ID, codersdk.UpdateMCPServerConfigRequest{
|
||||
OAuth2RevocationURL: &plaintextURL,
|
||||
})
|
||||
require.ErrorAs(t, err, &sdkErr)
|
||||
require.Equal(t, http.StatusBadRequest, sdkErr.StatusCode())
|
||||
|
||||
_, err = client.CreateMCPServerConfig(ctx, codersdk.CreateMCPServerConfigRequest{
|
||||
DisplayName: "Plaintext Revoke",
|
||||
Slug: "plaintext-revoke",
|
||||
Transport: "streamable_http",
|
||||
URL: "https://mcp.example.com/plaintext",
|
||||
AuthType: "oauth2",
|
||||
OAuth2ClientID: "cid",
|
||||
OAuth2AuthURL: "https://auth.example.com/authorize",
|
||||
OAuth2TokenURL: "https://auth.example.com/token",
|
||||
OAuth2RevocationURL: plaintextURL,
|
||||
Availability: "default_on",
|
||||
Enabled: true,
|
||||
ToolAllowList: []string{},
|
||||
ToolDenyList: []string{},
|
||||
})
|
||||
require.ErrorAs(t, err, &sdkErr)
|
||||
require.Equal(t, http.StatusBadRequest, sdkErr.StatusCode())
|
||||
|
||||
// An explicit empty string clears the stored URL.
|
||||
emptyURL := ""
|
||||
updated, err = client.UpdateMCPServerConfig(ctx, created.ID, codersdk.UpdateMCPServerConfigRequest{
|
||||
OAuth2RevocationURL: &emptyURL,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, updated.OAuth2RevocationURL)
|
||||
|
||||
updated, err = client.UpdateMCPServerConfig(ctx, created.ID, codersdk.UpdateMCPServerConfigRequest{
|
||||
OAuth2RevocationURL: &newRevocationURL,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, newRevocationURL, updated.OAuth2RevocationURL)
|
||||
|
||||
newAuth := "user_oidc"
|
||||
updated, err := client.UpdateMCPServerConfig(ctx, created.ID, codersdk.UpdateMCPServerConfigRequest{
|
||||
updated, err = client.UpdateMCPServerConfig(ctx, created.ID, codersdk.UpdateMCPServerConfigRequest{
|
||||
AuthType: &newAuth,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
@@ -404,6 +466,7 @@ func TestMCPServerConfigsUserOIDCClearsFields(t *testing.T) {
|
||||
require.Empty(t, updated.OAuth2ClientID)
|
||||
require.Empty(t, updated.OAuth2AuthURL)
|
||||
require.Empty(t, updated.OAuth2TokenURL)
|
||||
require.Empty(t, updated.OAuth2RevocationURL)
|
||||
require.Empty(t, updated.OAuth2Scopes)
|
||||
require.Empty(t, updated.APIKeyHeader)
|
||||
}
|
||||
@@ -525,65 +588,330 @@ func TestMCPServerConfigsUniqueSlug(t *testing.T) {
|
||||
func TestMCPServerConfigsOAuth2Disconnect(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)
|
||||
otherClient, other := coderdtest.CreateAnotherUser(t, adminClient, firstUser.OrganizationID)
|
||||
newDisconnectFixture := func(t *testing.T, slug, revocationURL string) (memberClient *codersdk.Client, memberID uuid.UUID, db database.Store, configID uuid.UUID) {
|
||||
t.Helper()
|
||||
|
||||
created, err := adminClient.CreateMCPServerConfig(ctx, codersdk.CreateMCPServerConfigRequest{
|
||||
DisplayName: "OAuth Disconnect Test",
|
||||
Slug: "oauth-disconnect",
|
||||
Transport: "streamable_http",
|
||||
URL: "https://mcp.example.com/oauth-disc",
|
||||
AuthType: "oauth2",
|
||||
OAuth2ClientID: "cid",
|
||||
OAuth2AuthURL: "https://auth.example.com/authorize",
|
||||
OAuth2TokenURL: "https://auth.example.com/token",
|
||||
Availability: "default_on",
|
||||
Enabled: true,
|
||||
ToolAllowList: []string{},
|
||||
ToolDenyList: []string{},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
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)
|
||||
|
||||
// Disconnect should succeed even when no token exists (idempotent).
|
||||
err = memberClient.MCPServerOAuth2Disconnect(ctx, created.ID)
|
||||
require.NoError(t, err)
|
||||
created, err := adminClient.CreateMCPServerConfig(ctx, codersdk.CreateMCPServerConfigRequest{
|
||||
DisplayName: "OAuth Disconnect " + slug,
|
||||
Slug: slug,
|
||||
Transport: "streamable_http",
|
||||
URL: "https://mcp.example.com/" + slug,
|
||||
AuthType: "oauth2",
|
||||
OAuth2ClientID: "cid",
|
||||
OAuth2AuthURL: "https://auth.example.com/authorize",
|
||||
OAuth2TokenURL: "https://auth.example.com/token",
|
||||
OAuth2RevocationURL: revocationURL,
|
||||
Availability: "default_on",
|
||||
Enabled: true,
|
||||
ToolAllowList: []string{},
|
||||
ToolDenyList: []string{},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
for _, userID := range []uuid.UUID{member.ID, other.ID} {
|
||||
return memberClient, member.ID, db, created.ID
|
||||
}
|
||||
|
||||
seedToken := func(t *testing.T, db database.Store, configID, userID uuid.UUID) {
|
||||
t.Helper()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
//nolint:gocritic // Seeding test state requires system access.
|
||||
_, err = db.UpsertMCPServerUserToken(dbauthz.AsSystemRestricted(ctx), database.UpsertMCPServerUserTokenParams{
|
||||
MCPServerConfigID: created.ID,
|
||||
_, err := db.UpsertMCPServerUserToken(dbauthz.AsSystemRestricted(ctx), database.UpsertMCPServerUserTokenParams{
|
||||
MCPServerConfigID: configID,
|
||||
UserID: userID,
|
||||
AccessToken: "valid-access",
|
||||
AccessToken: "access-token",
|
||||
RefreshToken: "refresh-token",
|
||||
TokenType: "Bearer",
|
||||
Expiry: sql.NullTime{Time: time.Now().Add(time.Hour), Valid: true},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
requireAuthConnected := func(client *codersdk.Client, want bool) {
|
||||
requireTokenDeleted := func(t *testing.T, db database.Store, configID, userID uuid.UUID) {
|
||||
t.Helper()
|
||||
configs, err := client.MCPServerConfigs(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, configs, 1)
|
||||
require.Equal(t, want, configs[0].AuthConnected)
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
//nolint:gocritic // Verifying persisted state requires system access.
|
||||
_, err := db.GetMCPServerUserToken(dbauthz.AsSystemRestricted(ctx), database.GetMCPServerUserTokenParams{
|
||||
MCPServerConfigID: configID,
|
||||
UserID: userID,
|
||||
})
|
||||
require.ErrorIs(t, err, sql.ErrNoRows)
|
||||
}
|
||||
requireAuthConnected(memberClient, true)
|
||||
requireAuthConnected(otherClient, true)
|
||||
|
||||
err = memberClient.MCPServerOAuth2Disconnect(ctx, created.ID)
|
||||
require.NoError(t, err)
|
||||
requireAuthConnected(memberClient, false)
|
||||
requireAuthConnected(otherClient, true)
|
||||
t.Run("NoToken", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
err = memberClient.MCPServerOAuth2Disconnect(ctx, created.ID)
|
||||
require.NoError(t, err)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
memberClient, _, _, configID := newDisconnectFixture(t, "disc-no-token", "")
|
||||
|
||||
resp, err := memberClient.MCPServerOAuth2DisconnectWithResponse(ctx, configID)
|
||||
require.NoError(t, err)
|
||||
require.False(t, resp.TokenRevoked)
|
||||
require.Empty(t, resp.TokenRevocationError)
|
||||
})
|
||||
|
||||
t.Run("DoesNotRevealHiddenConfigs", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
providerKeys := coderdtest.FakeOpenAICompatProviderAPIKeys(t)
|
||||
adminClient, _ := coderdtest.NewWithDatabase(t, &coderdtest.Options{
|
||||
DeploymentValues: mcpDeploymentValues(t),
|
||||
ChatProviderAPIKeys: &providerKeys,
|
||||
})
|
||||
firstUser := coderdtest.CreateFirstUser(t, adminClient)
|
||||
memberClient, _ := coderdtest.CreateAnotherUser(t, adminClient, firstUser.OrganizationID)
|
||||
|
||||
created, err := adminClient.CreateMCPServerConfig(ctx, codersdk.CreateMCPServerConfigRequest{
|
||||
DisplayName: "OAuth Disconnect Hidden",
|
||||
Slug: "disc-hidden",
|
||||
Transport: "streamable_http",
|
||||
URL: "https://mcp.example.com/disc-hidden",
|
||||
AuthType: "oauth2",
|
||||
OAuth2ClientID: "cid",
|
||||
OAuth2AuthURL: "https://auth.example.com/authorize",
|
||||
OAuth2TokenURL: "https://auth.example.com/token",
|
||||
Availability: "default_on",
|
||||
Enabled: false,
|
||||
ToolAllowList: []string{},
|
||||
ToolDenyList: []string{},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Disconnecting a disabled config the member cannot see must be
|
||||
// indistinguishable from disconnecting a nonexistent config ID.
|
||||
hiddenResp, err := memberClient.MCPServerOAuth2DisconnectWithResponse(ctx, created.ID)
|
||||
require.NoError(t, err)
|
||||
missingResp, err := memberClient.MCPServerOAuth2DisconnectWithResponse(ctx, uuid.New())
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, missingResp, hiddenResp)
|
||||
require.False(t, hiddenResp.TokenRevoked)
|
||||
require.Empty(t, hiddenResp.TokenRevocationError)
|
||||
})
|
||||
|
||||
t.Run("RevokesAtProvider", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
var gotForm atomic.Pointer[url.Values]
|
||||
revokeSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
require.NoError(t, r.ParseForm())
|
||||
form := r.PostForm
|
||||
gotForm.Store(&form)
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}))
|
||||
t.Cleanup(revokeSrv.Close)
|
||||
|
||||
memberClient, memberID, db, configID := newDisconnectFixture(t, "disc-revoke", revokeSrv.URL)
|
||||
seedToken(t, db, configID, memberID)
|
||||
|
||||
resp, err := memberClient.MCPServerOAuth2DisconnectWithResponse(ctx, configID)
|
||||
require.NoError(t, err)
|
||||
require.True(t, resp.TokenRevoked)
|
||||
require.Empty(t, resp.TokenRevocationError)
|
||||
|
||||
form := gotForm.Load()
|
||||
require.NotNil(t, form)
|
||||
require.Equal(t, "refresh-token", form.Get("token"))
|
||||
require.Equal(t, "refresh_token", form.Get("token_type_hint"))
|
||||
require.Equal(t, "cid", form.Get("client_id"))
|
||||
|
||||
requireTokenDeleted(t, db, configID, memberID)
|
||||
})
|
||||
|
||||
t.Run("RefreshCannotRestoreDisconnectedToken", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
refreshStarted := make(chan struct{})
|
||||
releaseRefresh := make(chan struct{})
|
||||
var releaseOnce sync.Once
|
||||
tokenSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
close(refreshStarted)
|
||||
select {
|
||||
case <-releaseRefresh:
|
||||
case <-r.Context().Done():
|
||||
return
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte(`{"access_token":"fresh-access","refresh_token":"fresh-refresh","token_type":"Bearer","expires_in":3600}`))
|
||||
}))
|
||||
t.Cleanup(tokenSrv.Close)
|
||||
t.Cleanup(func() { releaseOnce.Do(func() { close(releaseRefresh) }) })
|
||||
|
||||
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)
|
||||
|
||||
created, err := adminClient.CreateMCPServerConfig(ctx, codersdk.CreateMCPServerConfigRequest{
|
||||
DisplayName: "OAuth Disconnect Refresh Race",
|
||||
Slug: "disc-refresh-race",
|
||||
Transport: "streamable_http",
|
||||
URL: "https://mcp.example.com/disc-refresh-race",
|
||||
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: "old-refresh",
|
||||
TokenType: "Bearer",
|
||||
Expiry: sql.NullTime{Time: time.Now().Add(-time.Hour), Valid: true},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
type configResult struct {
|
||||
configs []codersdk.MCPServerConfig
|
||||
err error
|
||||
}
|
||||
result := make(chan configResult, 1)
|
||||
go func() {
|
||||
configs, listErr := memberClient.MCPServerConfigs(ctx)
|
||||
result <- configResult{configs: configs, err: listErr}
|
||||
}()
|
||||
|
||||
select {
|
||||
case <-refreshStarted:
|
||||
case <-ctx.Done():
|
||||
t.Fatal("timed out waiting for token refresh")
|
||||
}
|
||||
|
||||
_, err = memberClient.MCPServerOAuth2DisconnectWithResponse(ctx, created.ID)
|
||||
require.NoError(t, err)
|
||||
releaseOnce.Do(func() { close(releaseRefresh) })
|
||||
|
||||
var listed configResult
|
||||
select {
|
||||
case listed = <-result:
|
||||
case <-ctx.Done():
|
||||
t.Fatal("timed out waiting for refreshed config response")
|
||||
}
|
||||
require.NoError(t, listed.err)
|
||||
require.Len(t, listed.configs, 1)
|
||||
require.False(t, listed.configs[0].AuthConnected)
|
||||
requireTokenDeleted(t, db, created.ID, member.ID)
|
||||
})
|
||||
|
||||
t.Run("NoRevocationURL", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
memberClient, memberID, db, configID := newDisconnectFixture(t, "disc-no-url", "")
|
||||
seedToken(t, db, configID, memberID)
|
||||
|
||||
resp, err := memberClient.MCPServerOAuth2DisconnectWithResponse(ctx, configID)
|
||||
require.NoError(t, err)
|
||||
require.False(t, resp.TokenRevoked)
|
||||
require.Empty(t, resp.TokenRevocationError)
|
||||
|
||||
requireTokenDeleted(t, db, configID, memberID)
|
||||
})
|
||||
|
||||
t.Run("ProviderError", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
revokeSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
}))
|
||||
t.Cleanup(revokeSrv.Close)
|
||||
|
||||
memberClient, memberID, db, configID := newDisconnectFixture(t, "disc-err", revokeSrv.URL)
|
||||
seedToken(t, db, configID, memberID)
|
||||
|
||||
// Members get a generic error; provider bodies may echo the secret.
|
||||
resp, err := memberClient.MCPServerOAuth2DisconnectWithResponse(ctx, configID)
|
||||
require.NoError(t, err)
|
||||
require.False(t, resp.TokenRevoked)
|
||||
require.NotEmpty(t, resp.TokenRevocationError)
|
||||
require.NotContains(t, resp.TokenRevocationError, "HTTP 500")
|
||||
|
||||
requireTokenDeleted(t, db, configID, memberID)
|
||||
})
|
||||
|
||||
t.Run("OnlyDisconnectsCallingUser", func(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)
|
||||
otherClient, other := coderdtest.CreateAnotherUser(t, adminClient, firstUser.OrganizationID)
|
||||
|
||||
created, err := adminClient.CreateMCPServerConfig(ctx, codersdk.CreateMCPServerConfigRequest{
|
||||
DisplayName: "OAuth Disconnect Isolation",
|
||||
Slug: "disc-isolation",
|
||||
Transport: "streamable_http",
|
||||
URL: "https://mcp.example.com/disc-isolation",
|
||||
AuthType: "oauth2",
|
||||
OAuth2ClientID: "cid",
|
||||
OAuth2AuthURL: "https://auth.example.com/authorize",
|
||||
OAuth2TokenURL: "https://auth.example.com/token",
|
||||
Availability: "default_on",
|
||||
Enabled: true,
|
||||
ToolAllowList: []string{},
|
||||
ToolDenyList: []string{},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
for _, userID := range []uuid.UUID{member.ID, other.ID} {
|
||||
//nolint:gocritic // Seeding test state requires system access.
|
||||
_, err = db.UpsertMCPServerUserToken(dbauthz.AsSystemRestricted(ctx), database.UpsertMCPServerUserTokenParams{
|
||||
MCPServerConfigID: created.ID,
|
||||
UserID: userID,
|
||||
AccessToken: "valid-access",
|
||||
TokenType: "Bearer",
|
||||
Expiry: sql.NullTime{Time: time.Now().Add(time.Hour), Valid: true},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
requireAuthConnected := func(client *codersdk.Client, want bool) {
|
||||
t.Helper()
|
||||
configs, err := client.MCPServerConfigs(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, configs, 1)
|
||||
require.Equal(t, want, configs[0].AuthConnected)
|
||||
}
|
||||
requireAuthConnected(memberClient, true)
|
||||
requireAuthConnected(otherClient, true)
|
||||
|
||||
_, err = memberClient.MCPServerOAuth2DisconnectWithResponse(ctx, created.ID)
|
||||
require.NoError(t, err)
|
||||
requireAuthConnected(memberClient, false)
|
||||
requireAuthConnected(otherClient, true)
|
||||
|
||||
_, err = memberClient.MCPServerOAuth2DisconnectWithResponse(ctx, created.ID)
|
||||
require.NoError(t, err)
|
||||
})
|
||||
}
|
||||
|
||||
func TestMCPServerConfigsOAuth2AutoDiscovery(t *testing.T) {
|
||||
@@ -605,6 +933,7 @@ func TestMCPServerConfigsOAuth2AutoDiscovery(t *testing.T) {
|
||||
"authorization_endpoint": "` + "http://" + r.Host + `/authorize",
|
||||
"token_endpoint": "` + "http://" + r.Host + `/token",
|
||||
"registration_endpoint": "` + "http://" + r.Host + `/register",
|
||||
"revocation_endpoint": "` + "http://" + r.Host + `/revoke",
|
||||
"response_types_supported": ["code"],
|
||||
"scopes_supported": ["read", "write"]
|
||||
}`))
|
||||
@@ -665,7 +994,24 @@ func TestMCPServerConfigsOAuth2AutoDiscovery(t *testing.T) {
|
||||
require.True(t, created.HasOAuth2Secret)
|
||||
require.Equal(t, authServer.URL+"/authorize", created.OAuth2AuthURL)
|
||||
require.Equal(t, authServer.URL+"/token", created.OAuth2TokenURL)
|
||||
require.Equal(t, authServer.URL+"/revoke", created.OAuth2RevocationURL)
|
||||
require.Equal(t, "read write", created.OAuth2Scopes)
|
||||
|
||||
// An explicit revocation URL wins over the discovered one.
|
||||
overridden, err := client.CreateMCPServerConfig(ctx, codersdk.CreateMCPServerConfigRequest{
|
||||
DisplayName: "Auto-Discovery Override",
|
||||
Slug: "auto-discovery-override",
|
||||
Transport: "streamable_http",
|
||||
URL: mcpServer.URL + "/v1/mcp",
|
||||
AuthType: "oauth2",
|
||||
OAuth2RevocationURL: "https://override.example.com/revoke",
|
||||
Availability: "default_on",
|
||||
Enabled: true,
|
||||
ToolAllowList: []string{},
|
||||
ToolDenyList: []string{},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "https://override.example.com/revoke", overridden.OAuth2RevocationURL)
|
||||
})
|
||||
|
||||
// Verify that when both path-aware and root-level protected
|
||||
|
||||
+29
-4
@@ -4823,11 +4823,11 @@ func (p *Server) refreshMCPTokenIfNeeded(
|
||||
|
||||
//nolint:gocritic // Chatd needs system-level write access to
|
||||
// persist the refreshed OAuth2 token for the user.
|
||||
updated, err := p.db.UpsertMCPServerUserToken(
|
||||
updated, err := p.db.UpdateMCPServerUserTokenFromRefresh(
|
||||
dbauthz.AsSystemRestricted(ctx),
|
||||
database.UpsertMCPServerUserTokenParams{
|
||||
MCPServerConfigID: tok.MCPServerConfigID,
|
||||
UserID: tok.UserID,
|
||||
database.UpdateMCPServerUserTokenFromRefreshParams{
|
||||
ID: tok.ID,
|
||||
UpdatedAt: tok.UpdatedAt,
|
||||
AccessToken: result.AccessToken,
|
||||
AccessTokenKeyID: sql.NullString{},
|
||||
RefreshToken: result.RefreshToken,
|
||||
@@ -4837,6 +4837,31 @@ func (p *Server) refreshMCPTokenIfNeeded(
|
||||
},
|
||||
)
|
||||
if err != nil {
|
||||
if xerrors.Is(err, sql.ErrNoRows) {
|
||||
// A disconnect or re-authentication can win the optimistic update.
|
||||
//nolint:gocritic // Reading the winning token requires system access.
|
||||
current, readErr := p.db.GetMCPServerUserToken(
|
||||
dbauthz.AsSystemRestricted(ctx),
|
||||
database.GetMCPServerUserTokenParams{
|
||||
MCPServerConfigID: tok.MCPServerConfigID,
|
||||
UserID: tok.UserID,
|
||||
},
|
||||
)
|
||||
if readErr == nil {
|
||||
return current, nil
|
||||
}
|
||||
if !xerrors.Is(readErr, sql.ErrNoRows) {
|
||||
logger.Warn(ctx, "failed to load MCP oauth2 token after refresh conflict",
|
||||
slog.F("server_slug", cfg.Slug),
|
||||
slog.Error(readErr),
|
||||
)
|
||||
}
|
||||
tok.AccessToken = ""
|
||||
tok.RefreshToken = ""
|
||||
tok.Expiry = sql.NullTime{}
|
||||
return tok, nil
|
||||
}
|
||||
|
||||
// The provider may have rotated the refresh token,
|
||||
// invalidating the old one. Use the new token
|
||||
// in-memory so at least this connection succeeds.
|
||||
|
||||
@@ -184,6 +184,50 @@ func TestRefreshMCPTokenPermanentFailure(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func TestRefreshMCPTokenDeletedDuringRefresh(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tokenSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte(`{"access_token":"fresh-access","refresh_token":"fresh-refresh","token_type":"Bearer","expires_in":3600}`))
|
||||
}))
|
||||
t.Cleanup(tokenSrv.Close)
|
||||
|
||||
cfg := database.MCPServerConfig{
|
||||
ID: uuid.New(),
|
||||
Slug: "disconnected",
|
||||
AuthType: "oauth2",
|
||||
OAuth2ClientID: "cid",
|
||||
OAuth2TokenURL: tokenSrv.URL,
|
||||
}
|
||||
tok := expiredMCPToken(cfg.ID)
|
||||
|
||||
ctrl := gomock.NewController(t)
|
||||
db := dbmock.NewMockStore(ctrl)
|
||||
db.EXPECT().
|
||||
UpdateMCPServerUserTokenFromRefresh(gomock.Any(), gomock.Any()).
|
||||
DoAndReturn(func(_ context.Context, arg database.UpdateMCPServerUserTokenFromRefreshParams) (database.MCPServerUserToken, error) {
|
||||
require.Equal(t, tok.ID, arg.ID)
|
||||
require.Equal(t, tok.UpdatedAt, arg.UpdatedAt)
|
||||
return database.MCPServerUserToken{}, sql.ErrNoRows
|
||||
})
|
||||
db.EXPECT().
|
||||
GetMCPServerUserToken(gomock.Any(), database.GetMCPServerUserTokenParams{
|
||||
MCPServerConfigID: tok.MCPServerConfigID,
|
||||
UserID: tok.UserID,
|
||||
}).
|
||||
Return(database.MCPServerUserToken{}, sql.ErrNoRows)
|
||||
|
||||
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.Empty(t, result.OauthRefreshFailureReason)
|
||||
}
|
||||
|
||||
func TestRefreshExpiredMCPTokensSkipsFailedTokens(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
@@ -7,6 +7,8 @@ import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"slices"
|
||||
@@ -962,3 +964,211 @@ func RefreshOAuth2Token(
|
||||
Refreshed: refreshed,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// RevokeOAuth2Token revokes the user's token at the provider's RFC 7009
|
||||
// endpoint. It prefers the refresh token, retrying with the access token
|
||||
// only on unsupported_token_type; other failures do not fall back, since
|
||||
// an access-token success would hide a possibly live refresh token.
|
||||
// Returns false without error when there is no revocation endpoint or no
|
||||
// stored token. Errors carry only the HTTP status because provider
|
||||
// bodies may echo secrets.
|
||||
func RevokeOAuth2Token(
|
||||
ctx context.Context,
|
||||
httpClient *http.Client,
|
||||
cfg database.MCPServerConfig,
|
||||
tok database.MCPServerUserToken,
|
||||
) (bool, error) {
|
||||
if cfg.OAuth2RevocationURL == "" {
|
||||
return false, nil
|
||||
}
|
||||
if tok.RefreshToken == "" && tok.AccessToken == "" {
|
||||
return false, nil
|
||||
}
|
||||
if err := ValidateRevocationEndpoint(cfg.OAuth2RevocationURL); err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
if httpClient == nil {
|
||||
httpClient = mcpHTTPClient()
|
||||
}
|
||||
if httpClient == nil {
|
||||
httpClient = http.DefaultClient
|
||||
}
|
||||
// Copy so CheckRedirect does not leak into the shared client.
|
||||
redirectSafe := *httpClient
|
||||
redirectSafe.CheckRedirect = checkRevocationRedirect
|
||||
httpClient = &redirectSafe
|
||||
|
||||
token, hint := tok.AccessToken, "access_token"
|
||||
if tok.RefreshToken != "" {
|
||||
token, hint = tok.RefreshToken, "refresh_token"
|
||||
}
|
||||
status, errorCode, err := postTokenRevocation(ctx, httpClient, cfg, token, hint)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
if isRevocationSuccessStatus(status) {
|
||||
return true, nil
|
||||
}
|
||||
|
||||
if hint == "refresh_token" && tok.AccessToken != "" && errorCode == "unsupported_token_type" {
|
||||
fbStatus, _, fbErr := postTokenRevocation(ctx, httpClient, cfg, tok.AccessToken, "access_token")
|
||||
if fbErr != nil {
|
||||
return false, fbErr
|
||||
}
|
||||
if isRevocationSuccessStatus(fbStatus) {
|
||||
return true, nil
|
||||
}
|
||||
return false, xerrors.Errorf(
|
||||
"revocation endpoint returned HTTP %d for the refresh token and HTTP %d for the access token",
|
||||
status, fbStatus,
|
||||
)
|
||||
}
|
||||
return false, xerrors.Errorf(
|
||||
"revocation endpoint returned HTTP %d", status,
|
||||
)
|
||||
}
|
||||
|
||||
func isLoopbackHost(host string) bool {
|
||||
if strings.EqualFold(host, "localhost") {
|
||||
return true
|
||||
}
|
||||
ip := net.ParseIP(host)
|
||||
return ip != nil && ip.IsLoopback()
|
||||
}
|
||||
|
||||
// ValidateRevocationEndpoint enforces the RFC 7009 HTTPS requirement;
|
||||
// the request carries token material and the client secret. Plain HTTP
|
||||
// is allowed only for loopback hosts.
|
||||
func ValidateRevocationEndpoint(rawURL string) error {
|
||||
parsed, err := url.Parse(rawURL)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("parse revocation URL: %w", err)
|
||||
}
|
||||
// url.Parse accepts hostless forms like "https:/revoke" that can
|
||||
// never be POSTed to.
|
||||
if parsed.Hostname() == "" {
|
||||
return xerrors.Errorf(
|
||||
"revocation endpoint %q has no host", parsed.Redacted(),
|
||||
)
|
||||
}
|
||||
if !isAllowedRevocationScheme(parsed) {
|
||||
return xerrors.Errorf(
|
||||
"revocation endpoint %q must use https", parsed.Redacted(),
|
||||
)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func isAllowedRevocationScheme(u *url.URL) bool {
|
||||
if u.Scheme == "https" {
|
||||
return true
|
||||
}
|
||||
return u.Scheme == "http" && isLoopbackHost(u.Hostname())
|
||||
}
|
||||
|
||||
// checkRevocationRedirect stops the revocation POST, which carries
|
||||
// token material and client credentials, from following redirects off
|
||||
// the provider's origin. Loopback to loopback is exempt.
|
||||
func checkRevocationRedirect(req *http.Request, via []*http.Request) error {
|
||||
if len(via) >= 10 {
|
||||
return xerrors.New("stopped after 10 redirects")
|
||||
}
|
||||
// net/http follows 301/302/303 with a bodyless GET; the token never
|
||||
// reaches the endpoint and a trailing 200 would be a false success.
|
||||
if req.Method != http.MethodPost {
|
||||
return xerrors.New(
|
||||
"revocation redirect dropped the POST body",
|
||||
)
|
||||
}
|
||||
if !isAllowedRevocationScheme(req.URL) {
|
||||
return xerrors.New("revocation redirect target must use https")
|
||||
}
|
||||
origin := via[0].URL
|
||||
if isLoopbackHost(req.URL.Hostname()) && isLoopbackHost(origin.Hostname()) {
|
||||
return nil
|
||||
}
|
||||
if req.URL.Scheme != origin.Scheme ||
|
||||
!strings.EqualFold(req.URL.Hostname(), origin.Hostname()) ||
|
||||
normalizedPort(req.URL) != normalizedPort(origin) {
|
||||
return xerrors.Errorf(
|
||||
"revocation redirect must stay on origin %q",
|
||||
origin.Scheme+"://"+origin.Host,
|
||||
)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func normalizedPort(u *url.URL) string {
|
||||
if p := u.Port(); p != "" {
|
||||
return p
|
||||
}
|
||||
switch u.Scheme {
|
||||
case "https":
|
||||
return "443"
|
||||
case "http":
|
||||
return "80"
|
||||
default:
|
||||
return ""
|
||||
}
|
||||
}
|
||||
|
||||
func isRevocationSuccessStatus(status int) bool {
|
||||
return status == http.StatusOK || status == http.StatusNoContent
|
||||
}
|
||||
|
||||
// postTokenRevocation returns the HTTP status and the RFC 6749 error
|
||||
// code from the body; the raw body never propagates.
|
||||
func postTokenRevocation(
|
||||
ctx context.Context,
|
||||
httpClient *http.Client,
|
||||
cfg database.MCPServerConfig,
|
||||
token, tokenTypeHint string,
|
||||
) (int, string, error) {
|
||||
form := url.Values{}
|
||||
form.Set("token", token)
|
||||
form.Set("token_type_hint", tokenTypeHint)
|
||||
// Only public clients send client_id in the body; mixing it with
|
||||
// Basic auth is malformed per RFC 6749 section 2.3.1.
|
||||
if cfg.OAuth2ClientSecret == "" {
|
||||
form.Set("client_id", cfg.OAuth2ClientID)
|
||||
}
|
||||
|
||||
revokeCtx, cancel := context.WithTimeout(ctx, connectTimeout)
|
||||
defer cancel()
|
||||
|
||||
req, err := http.NewRequestWithContext(
|
||||
revokeCtx, http.MethodPost,
|
||||
cfg.OAuth2RevocationURL, strings.NewReader(form.Encode()),
|
||||
)
|
||||
if err != nil {
|
||||
return 0, "", xerrors.Errorf("create revocation request: %w", err)
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
||||
// Credentials are form-encoded per RFC 6749 section 2.3.1
|
||||
// (mirrors x/oauth2).
|
||||
if cfg.OAuth2ClientSecret != "" {
|
||||
req.SetBasicAuth(url.QueryEscape(cfg.OAuth2ClientID), url.QueryEscape(cfg.OAuth2ClientSecret))
|
||||
}
|
||||
|
||||
resp, err := httpClient.Do(req)
|
||||
if err != nil {
|
||||
var urlErr *url.Error
|
||||
if errors.As(err, &urlErr) {
|
||||
err = urlErr.Err
|
||||
}
|
||||
return 0, "", xerrors.Errorf("revoke oauth2 token: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if isRevocationSuccessStatus(resp.StatusCode) {
|
||||
_, _ = io.Copy(io.Discard, resp.Body)
|
||||
return resp.StatusCode, "", nil
|
||||
}
|
||||
var errBody struct {
|
||||
Error string `json:"error"`
|
||||
}
|
||||
_ = json.NewDecoder(io.LimitReader(resp.Body, 4096)).Decode(&errBody)
|
||||
_, _ = io.Copy(io.Discard, resp.Body)
|
||||
return resp.StatusCode, errBody.Error, nil
|
||||
}
|
||||
|
||||
@@ -0,0 +1,89 @@
|
||||
package mcpclient
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"net/url"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestCheckRevocationRedirect(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
req := func(method, rawURL string) *http.Request {
|
||||
u, err := url.Parse(rawURL)
|
||||
require.NoError(t, err)
|
||||
return &http.Request{Method: method, URL: u}
|
||||
}
|
||||
|
||||
origin := "https://provider.example/revoke"
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
req *http.Request
|
||||
origin string
|
||||
wantErr string
|
||||
wantAbsent string
|
||||
}{
|
||||
{
|
||||
name: "SamePathOnOrigin",
|
||||
req: req(http.MethodPost, "https://provider.example/revoke2"),
|
||||
},
|
||||
{
|
||||
name: "ExplicitDefaultPort",
|
||||
req: req(http.MethodPost, "https://provider.example:443/revoke2"),
|
||||
},
|
||||
{
|
||||
name: "DifferentPort",
|
||||
req: req(http.MethodPost, "https://provider.example:8443/collect"),
|
||||
wantErr: "must stay on origin",
|
||||
},
|
||||
{
|
||||
name: "DifferentHost",
|
||||
req: req(http.MethodPost, "https://attacker.example/collect?token=reflected-token#fragment"),
|
||||
wantErr: "must stay on origin",
|
||||
wantAbsent: "reflected-token",
|
||||
},
|
||||
{
|
||||
name: "BodyDroppingGet",
|
||||
req: req(http.MethodGet, "https://provider.example/other"),
|
||||
wantErr: "dropped the POST body",
|
||||
},
|
||||
{
|
||||
name: "PlaintextTarget",
|
||||
req: req(http.MethodPost, "http://provider.example/revoke"),
|
||||
wantErr: "must use https",
|
||||
},
|
||||
{
|
||||
name: "LoopbackToLoopbackAnyPort",
|
||||
req: req(http.MethodPost, "http://127.0.0.1:9999/revoke"),
|
||||
origin: "http://localhost:1234/revoke",
|
||||
},
|
||||
{
|
||||
name: "OriginToLoopback",
|
||||
req: req(http.MethodPost, "http://localhost:1234/revoke"),
|
||||
wantErr: "must stay on origin",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
o := tc.origin
|
||||
if o == "" {
|
||||
o = origin
|
||||
}
|
||||
err := checkRevocationRedirect(tc.req, []*http.Request{req(http.MethodPost, o)})
|
||||
if tc.wantErr == "" {
|
||||
require.NoError(t, err)
|
||||
return
|
||||
}
|
||||
require.ErrorContains(t, err, tc.wantErr)
|
||||
if tc.wantAbsent != "" {
|
||||
require.NotContains(t, err.Error(), tc.wantAbsent)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,413 @@
|
||||
package mcpclient_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/mcpclient"
|
||||
)
|
||||
|
||||
type revokeRequest struct {
|
||||
form map[string][]string
|
||||
basicUser string
|
||||
basicPass string
|
||||
basicSet bool
|
||||
}
|
||||
|
||||
func captureRevoke(t *testing.T, got chan<- revokeRequest) http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
require.NoError(t, r.ParseForm())
|
||||
user, pass, ok := r.BasicAuth()
|
||||
got <- revokeRequest{form: r.PostForm, basicUser: user, basicPass: pass, basicSet: ok}
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRevokeOAuth2Token(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
t.Run("NoRevocationURL", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
revoked, err := mcpclient.RevokeOAuth2Token(
|
||||
context.Background(),
|
||||
nil,
|
||||
database.MCPServerConfig{OAuth2ClientID: "cid"},
|
||||
database.MCPServerUserToken{AccessToken: "at", RefreshToken: "rt"},
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.False(t, revoked)
|
||||
})
|
||||
|
||||
t.Run("RevokesRefreshToken", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
got := make(chan revokeRequest, 1)
|
||||
srv := httptest.NewServer(captureRevoke(t, got))
|
||||
defer srv.Close()
|
||||
|
||||
revoked, err := mcpclient.RevokeOAuth2Token(
|
||||
context.Background(),
|
||||
srv.Client(),
|
||||
database.MCPServerConfig{
|
||||
OAuth2ClientID: "cid",
|
||||
OAuth2RevocationURL: srv.URL,
|
||||
},
|
||||
database.MCPServerUserToken{AccessToken: "at", RefreshToken: "rt"},
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.True(t, revoked)
|
||||
c := <-got
|
||||
require.Equal(t, []string{"rt"}, c.form["token"])
|
||||
require.Equal(t, []string{"refresh_token"}, c.form["token_type_hint"])
|
||||
require.Equal(t, []string{"cid"}, c.form["client_id"])
|
||||
// Public clients must not authenticate.
|
||||
require.False(t, c.basicSet)
|
||||
require.NotContains(t, c.form, "client_secret")
|
||||
})
|
||||
|
||||
t.Run("NoContentIsSuccess", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.WriteHeader(http.StatusNoContent)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
revoked, err := mcpclient.RevokeOAuth2Token(
|
||||
context.Background(),
|
||||
srv.Client(),
|
||||
database.MCPServerConfig{
|
||||
OAuth2ClientID: "cid",
|
||||
OAuth2RevocationURL: srv.URL,
|
||||
},
|
||||
database.MCPServerUserToken{AccessToken: "at"},
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.True(t, revoked)
|
||||
})
|
||||
|
||||
t.Run("AcceptedIsNotSuccess", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.WriteHeader(http.StatusAccepted)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
revoked, err := mcpclient.RevokeOAuth2Token(
|
||||
context.Background(),
|
||||
srv.Client(),
|
||||
database.MCPServerConfig{
|
||||
OAuth2ClientID: "cid",
|
||||
OAuth2RevocationURL: srv.URL,
|
||||
},
|
||||
database.MCPServerUserToken{AccessToken: "at"},
|
||||
)
|
||||
require.ErrorContains(t, err, "HTTP 202")
|
||||
require.False(t, revoked)
|
||||
})
|
||||
|
||||
t.Run("AccessTokenFallbackWithBasicAuth", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
got := make(chan revokeRequest, 1)
|
||||
srv := httptest.NewServer(captureRevoke(t, got))
|
||||
defer srv.Close()
|
||||
|
||||
revoked, err := mcpclient.RevokeOAuth2Token(
|
||||
context.Background(),
|
||||
srv.Client(),
|
||||
database.MCPServerConfig{
|
||||
OAuth2ClientID: "cid",
|
||||
OAuth2ClientSecret: "secret",
|
||||
OAuth2RevocationURL: srv.URL,
|
||||
},
|
||||
database.MCPServerUserToken{AccessToken: "at"},
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.True(t, revoked)
|
||||
c := <-got
|
||||
require.Equal(t, []string{"at"}, c.form["token"])
|
||||
require.Equal(t, []string{"access_token"}, c.form["token_type_hint"])
|
||||
// Basic auth must not be mixed with body client_id (RFC 6749 2.3.1).
|
||||
require.True(t, c.basicSet)
|
||||
require.Equal(t, "cid", c.basicUser)
|
||||
require.Equal(t, "secret", c.basicPass)
|
||||
require.NotContains(t, c.form, "client_id")
|
||||
require.NotContains(t, c.form, "client_secret")
|
||||
})
|
||||
|
||||
t.Run("AccessTokenFallbackAfterUnsupportedTokenType", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
got := make(chan revokeRequest, 2)
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
require.NoError(t, r.ParseForm())
|
||||
got <- revokeRequest{form: r.PostForm}
|
||||
if r.PostForm.Get("token_type_hint") == "refresh_token" {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
_, _ = w.Write([]byte(`{"error":"unsupported_token_type"}`))
|
||||
return
|
||||
}
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
revoked, err := mcpclient.RevokeOAuth2Token(
|
||||
context.Background(),
|
||||
srv.Client(),
|
||||
database.MCPServerConfig{
|
||||
OAuth2ClientID: "cid",
|
||||
OAuth2RevocationURL: srv.URL,
|
||||
},
|
||||
database.MCPServerUserToken{AccessToken: "at", RefreshToken: "rt"},
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.True(t, revoked)
|
||||
first := <-got
|
||||
require.Equal(t, []string{"rt"}, first.form["token"])
|
||||
require.Equal(t, []string{"refresh_token"}, first.form["token_type_hint"])
|
||||
second := <-got
|
||||
require.Equal(t, []string{"at"}, second.form["token"])
|
||||
require.Equal(t, []string{"access_token"}, second.form["token_type_hint"])
|
||||
})
|
||||
|
||||
t.Run("NoFallbackWithoutUnsupportedTokenType", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var calls atomic.Int64
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
calls.Add(1)
|
||||
w.WriteHeader(http.StatusUnauthorized)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
revoked, err := mcpclient.RevokeOAuth2Token(
|
||||
context.Background(),
|
||||
srv.Client(),
|
||||
database.MCPServerConfig{
|
||||
OAuth2ClientID: "cid",
|
||||
OAuth2RevocationURL: srv.URL,
|
||||
},
|
||||
database.MCPServerUserToken{AccessToken: "at", RefreshToken: "rt"},
|
||||
)
|
||||
require.Error(t, err)
|
||||
require.False(t, revoked)
|
||||
require.Contains(t, err.Error(), "HTTP 401")
|
||||
// No access-token fallback: it could mask a live refresh token.
|
||||
require.EqualValues(t, 1, calls.Load())
|
||||
})
|
||||
|
||||
t.Run("FallbackAlsoFails", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
require.NoError(t, r.ParseForm())
|
||||
if r.PostForm.Get("token_type_hint") == "refresh_token" {
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
_, _ = w.Write([]byte(`{"error":"unsupported_token_type"}`))
|
||||
return
|
||||
}
|
||||
w.WriteHeader(http.StatusServiceUnavailable)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
revoked, err := mcpclient.RevokeOAuth2Token(
|
||||
context.Background(),
|
||||
srv.Client(),
|
||||
database.MCPServerConfig{
|
||||
OAuth2ClientID: "cid",
|
||||
OAuth2RevocationURL: srv.URL,
|
||||
},
|
||||
database.MCPServerUserToken{AccessToken: "at", RefreshToken: "rt"},
|
||||
)
|
||||
require.Error(t, err)
|
||||
require.False(t, revoked)
|
||||
require.Contains(t, err.Error(), "HTTP 400 for the refresh token")
|
||||
require.Contains(t, err.Error(), "HTTP 503 for the access token")
|
||||
})
|
||||
|
||||
t.Run("RejectsNonHTTPSEndpoint", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// Loopback is exempt only for plain http; hostless forms
|
||||
// parse but can never be POSTed to.
|
||||
for u, wantErr := range map[string]string{
|
||||
"http://revoke.example.com/revoke": "must use https",
|
||||
"ftp://localhost/revoke": "must use https",
|
||||
"https:/revoke": "has no host",
|
||||
"https:///revoke": "has no host",
|
||||
} {
|
||||
revoked, err := mcpclient.RevokeOAuth2Token(
|
||||
context.Background(),
|
||||
nil,
|
||||
database.MCPServerConfig{
|
||||
OAuth2ClientID: "cid",
|
||||
OAuth2RevocationURL: u,
|
||||
},
|
||||
database.MCPServerUserToken{AccessToken: "at", RefreshToken: "rt"},
|
||||
)
|
||||
require.Error(t, err, u)
|
||||
require.False(t, revoked, u)
|
||||
require.Contains(t, err.Error(), wantErr, u)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("RejectsPlaintextRedirect", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
http.Redirect(w, r, "http://revoke.example.com/revoke", http.StatusTemporaryRedirect)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
revoked, err := mcpclient.RevokeOAuth2Token(
|
||||
context.Background(),
|
||||
srv.Client(),
|
||||
database.MCPServerConfig{
|
||||
OAuth2ClientID: "cid",
|
||||
OAuth2RevocationURL: srv.URL,
|
||||
},
|
||||
database.MCPServerUserToken{AccessToken: "at", RefreshToken: "rt"},
|
||||
)
|
||||
require.Error(t, err)
|
||||
require.False(t, revoked)
|
||||
require.Contains(t, err.Error(), "must use https")
|
||||
})
|
||||
|
||||
t.Run("RejectsBodyDroppingRedirect", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
target := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
// Returns 200 to the bodyless GET produced by the redirect.
|
||||
require.NoError(t, r.ParseForm())
|
||||
require.Empty(t, r.PostForm.Get("token"))
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}))
|
||||
defer target.Close()
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
http.Redirect(w, r, target.URL, http.StatusFound)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
revoked, err := mcpclient.RevokeOAuth2Token(
|
||||
context.Background(),
|
||||
srv.Client(),
|
||||
database.MCPServerConfig{
|
||||
OAuth2ClientID: "cid",
|
||||
OAuth2RevocationURL: srv.URL,
|
||||
},
|
||||
database.MCPServerUserToken{AccessToken: "at", RefreshToken: "rt"},
|
||||
)
|
||||
require.Error(t, err)
|
||||
require.False(t, revoked)
|
||||
require.Contains(t, err.Error(), "dropped the POST body")
|
||||
})
|
||||
|
||||
t.Run("RejectsCrossHostRedirect", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// CheckRedirect rejects before the attacker host is dialed.
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
http.Redirect(w, r, "https://attacker.example.com/collect?token=reflected-token", http.StatusTemporaryRedirect)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
revoked, err := mcpclient.RevokeOAuth2Token(
|
||||
context.Background(),
|
||||
srv.Client(),
|
||||
database.MCPServerConfig{
|
||||
OAuth2ClientID: "cid",
|
||||
OAuth2RevocationURL: srv.URL,
|
||||
},
|
||||
database.MCPServerUserToken{AccessToken: "at", RefreshToken: "rt"},
|
||||
)
|
||||
require.Error(t, err)
|
||||
require.False(t, revoked)
|
||||
require.Contains(t, err.Error(), "must stay on origin")
|
||||
require.NotContains(t, err.Error(), "/collect")
|
||||
require.NotContains(t, err.Error(), "reflected-token")
|
||||
})
|
||||
|
||||
t.Run("FollowsLoopbackRedirect", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
got := make(chan revokeRequest, 1)
|
||||
target := httptest.NewServer(captureRevoke(t, got))
|
||||
defer target.Close()
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
http.Redirect(w, r, target.URL, http.StatusTemporaryRedirect)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
revoked, err := mcpclient.RevokeOAuth2Token(
|
||||
context.Background(),
|
||||
srv.Client(),
|
||||
database.MCPServerConfig{
|
||||
OAuth2ClientID: "cid",
|
||||
OAuth2RevocationURL: srv.URL,
|
||||
},
|
||||
database.MCPServerUserToken{AccessToken: "at", RefreshToken: "rt"},
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.True(t, revoked)
|
||||
c := <-got
|
||||
require.Equal(t, []string{"rt"}, c.form["token"])
|
||||
})
|
||||
|
||||
t.Run("NoTokenMaterial", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
t.Error("provider must not be called without token material")
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
revoked, err := mcpclient.RevokeOAuth2Token(
|
||||
context.Background(),
|
||||
srv.Client(),
|
||||
database.MCPServerConfig{
|
||||
OAuth2ClientID: "cid",
|
||||
OAuth2RevocationURL: srv.URL,
|
||||
},
|
||||
database.MCPServerUserToken{},
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.False(t, revoked)
|
||||
})
|
||||
|
||||
t.Run("ProviderError", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
_, _ = w.Write([]byte("SECRET-ECHO " + strings.Repeat("x", 2048)))
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
revoked, err := mcpclient.RevokeOAuth2Token(
|
||||
context.Background(),
|
||||
srv.Client(),
|
||||
database.MCPServerConfig{
|
||||
OAuth2ClientID: "cid",
|
||||
OAuth2RevocationURL: srv.URL,
|
||||
},
|
||||
database.MCPServerUserToken{AccessToken: "at", RefreshToken: "rt"},
|
||||
)
|
||||
require.Error(t, err)
|
||||
require.False(t, revoked)
|
||||
require.Contains(t, err.Error(), "HTTP 500")
|
||||
// The secret-echoing body must not surface in the error.
|
||||
require.NotContains(t, err.Error(), "SECRET-ECHO")
|
||||
})
|
||||
}
|
||||
Reference in New Issue
Block a user