feat!: org-scope MCP server configs with RBAC (#27942)

Moves `mcp_server_configs` from deployment scope to organization scope
so each organization fully controls the MCP servers its members can use
with Coder Agents.

## Summary

- Migration: adds `organization_id` (NOT NULL, FK) and keeps existing
rows as the default organization's originals with credentials intact.
Other organizations start with no MCP servers and configure their own;
nothing is copied across organizations. Chats outside the default
organization keep any now-cross-organization `mcp_server_ids` entries;
the runtime already ignores IDs that do not resolve in the chat's
organization, so no data rewrite is needed. Slug uniqueness becomes
`(organization_id, slug)`.
- RBAC: new org-scoped `ResourceMCPServerConfig` with regosql converter
and `GetAuthorizedMCPServerConfigs`; org admins get in-org CRUD, org
members get read (replaced by ACL evaluation in the follow-up ACL PR in
this stack).
- API: all config routes nest under the organization, matching
templates: `POST|GET
/api/experimental/organizations/{organization}/mcp-servers` and
`GET|PATCH|DELETE .../mcp-servers/{mcpserverconfig}` (plus
`oauth2/connect`), resolved by a read-only param middleware that
conceals read-denied and cross-organization access as 404. Two routes
stay on the frozen `/api/experimental/mcp/servers/{mcpServer}` block:
the OAuth2 callback (the redirect URI baked into existing AS-side client
registrations) and `oauth2/disconnect`, which must remain reachable by
users removed from the organization so they can still revoke their token
grant.
- Chat runtime: selection validation and generation resolve configs
strictly by IDs, enabled state, and the chat's organization in SQL;
requested duplicates are normalized; invalid or cross-org IDs are
rejected with the precise ID list. IDs already persisted on a chat are
exempt from message-time rejection so disabling a selected server never
blocks sends; generation skips servers that are no longer usable.
- Frontend: API layer and admin settings pages target the new endpoints.
The admin page manages the default organization's servers; the org
picker is tracked separately (CODAGT-714).

- Security hardening from review: OAuth user grants are additionally
bound to `oauth2_revocation_url` (changing it invalidates grants, and a
racing OAuth callback gets 409 instead of recreating a grant).
Stack-wide SSRF protection for MCP config-directed traffic was split
into its own PR at the top of this stack (#28242) to keep this diff
reviewable; this PR keeps main's existing discovery IP-range guard.
OAuth2 auto-discovery now completes before the config row is inserted: a
failed discovery persists nothing, and there is no provisional row that
concurrent updates could race against.

Two follow-up PRs in this stack were split out to keep this diff
reviewable: #28064 completes the swagger annotations for the moved
routes (main already ships these experimental MCP handlers unannotated),
and #28065 carries hardening fixes and regression pins on top of the
cutover.

- Force On enforcement (landed on main mid-review) is org-scoped: the
forced set is read per chat organization
(`GetForcedMCPServerConfigsByOrganization`), so another organization's
`force_on` server never attaches to a chat.

## Breaking changes (experimental API)

The MCP server config endpoints move from the deployment-scoped
`/api/experimental/mcp/servers` block to organization-nested paths:
`POST|GET /api/experimental/organizations/{organization}/mcp-servers`
and `GET|PATCH|DELETE
/api/experimental/organizations/{organization}/mcp-servers/{mcpserverconfig}`
(plus `oauth2/connect`). The old paths are removed, so API consumers
must supply an organization. Two routes intentionally stay on the frozen
`/api/experimental/mcp/servers/{mcpServer}` block: the OAuth2 callback
(its redirect URI is baked into existing AS-side client registrations)
and `oauth2/disconnect` (must remain reachable by users removed from the
organization). These endpoints are under `/api/experimental`, so no
deprecation window is provided.

## Rolling upgrades

During a rolling deploy, an old replica creating an MCP config can fail
the new `NOT NULL organization_id` constraint until it is upgraded
(reads are unaffected: old binaries' generated queries select their own
column lists). This matches the repo's existing precedent for additive
NOT NULL migrations (000562) and affects only the admin config-create
path in the upgrade window.

Upgrades are expected to run in scheduled maintenance downtime with the
database locked during migration, so the migration ships no
rolling-upgrade compatibility machinery. The down migration deletes
organization-created configs (their chat references are cleaned by the
000510 delete trigger) and restores deployment-wide slug uniqueness.

Part of the MCP org-separation stack (CODAGT-711 -> CODAGT-717 audit ->
CODAGT-712 ACLs -> CODAGT-806 token RBAC).

Closes https://linear.app/codercom/issue/CODAGT-711

UAT: validated end to end on a two-org dogfood deployment, including a
real pre-migration to post-migration upgrade, cross-org isolation
(404s), same-slug-two-orgs, chat selection gating, and a live MCP tool
call through the org-scoped generation path. The migration was later
revised to keep existing rows in the default organization only (no
per-organization copies); that revision is covered by the migration test
suite.

> Mux (AI agent) authored this PR on Mike's behalf.

<!-- mux-attribution: model=claude-fable-5 thinking=high -->

---------

Co-authored-by: Mathias Fredriksson <mafredri@gmail.com>
This commit is contained in:
Michael Suchacz
2026-08-19 18:04:13 +00:00
committed by GitHub
co-authored by Mathias Fredriksson
parent 7ff1278ab3
commit 443e3b9b80
72 changed files with 3649 additions and 988 deletions
+12
View File
@@ -16509,6 +16509,11 @@ const docTemplate = `{
"license:create",
"license:delete",
"license:read",
"mcp_server_config:*",
"mcp_server_config:create",
"mcp_server_config:delete",
"mcp_server_config:read",
"mcp_server_config:update",
"notification_message:*",
"notification_message:create",
"notification_message:delete",
@@ -16749,6 +16754,11 @@ const docTemplate = `{
"APIKeyScopeLicenseCreate",
"APIKeyScopeLicenseDelete",
"APIKeyScopeLicenseRead",
"APIKeyScopeMcpServerConfigAll",
"APIKeyScopeMcpServerConfigCreate",
"APIKeyScopeMcpServerConfigDelete",
"APIKeyScopeMcpServerConfigRead",
"APIKeyScopeMcpServerConfigUpdate",
"APIKeyScopeNotificationMessageAll",
"APIKeyScopeNotificationMessageCreate",
"APIKeyScopeNotificationMessageDelete",
@@ -24036,6 +24046,7 @@ const docTemplate = `{
"idpsync_settings",
"inbox_notification",
"license",
"mcp_server_config",
"notification_message",
"notification_preference",
"notification_template",
@@ -24089,6 +24100,7 @@ const docTemplate = `{
"ResourceIdpsyncSettings",
"ResourceInboxNotification",
"ResourceLicense",
"ResourceMCPServerConfig",
"ResourceNotificationMessage",
"ResourceNotificationPreference",
"ResourceNotificationTemplate",
+12
View File
@@ -14774,6 +14774,11 @@
"license:create",
"license:delete",
"license:read",
"mcp_server_config:*",
"mcp_server_config:create",
"mcp_server_config:delete",
"mcp_server_config:read",
"mcp_server_config:update",
"notification_message:*",
"notification_message:create",
"notification_message:delete",
@@ -15014,6 +15019,11 @@
"APIKeyScopeLicenseCreate",
"APIKeyScopeLicenseDelete",
"APIKeyScopeLicenseRead",
"APIKeyScopeMcpServerConfigAll",
"APIKeyScopeMcpServerConfigCreate",
"APIKeyScopeMcpServerConfigDelete",
"APIKeyScopeMcpServerConfigRead",
"APIKeyScopeMcpServerConfigUpdate",
"APIKeyScopeNotificationMessageAll",
"APIKeyScopeNotificationMessageCreate",
"APIKeyScopeNotificationMessageDelete",
@@ -22028,6 +22038,7 @@
"idpsync_settings",
"inbox_notification",
"license",
"mcp_server_config",
"notification_message",
"notification_preference",
"notification_template",
@@ -22081,6 +22092,7 @@
"ResourceIdpsyncSettings",
"ResourceInboxNotification",
"ResourceLicense",
"ResourceMCPServerConfig",
"ResourceNotificationMessage",
"ResourceNotificationPreference",
"ResourceNotificationTemplate",
+22 -14
View File
@@ -1378,6 +1378,23 @@ func New(options *Options) *API {
r.Use(httpmw.RateLimit(options.FilesRateLimit, time.Minute))
r.Get("/chats/files/{file}/download", api.downloadChatFile)
})
r.Route("/organizations", func(r chi.Router) {
r.Use(apiKeyMiddleware)
r.Route("/{organization}", func(r chi.Router) {
r.Use(httpmw.ExtractOrganizationParam(options.Database))
r.Route("/mcp-servers", func(r chi.Router) {
r.Get("/", api.listMCPServerConfigs)
r.Post("/", api.createMCPServerConfig)
r.Route("/{mcpserverconfig}", func(r chi.Router) {
r.Use(httpmw.ExtractMCPServerConfigParam(options.Database))
r.Get("/", api.getMCPServerConfig)
r.Patch("/", api.updateMCPServerConfig)
r.Delete("/", api.deleteMCPServerConfig)
r.Get("/oauth2/connect", api.mcpServerOAuth2Connect)
})
})
})
})
r.Route("/chats", func(r chi.Router) {
r.Use(
apiKeyMiddleware,
@@ -1499,20 +1516,11 @@ func New(options *Options) *API {
r.Use(
apiKeyMiddleware,
)
// MCP server configuration endpoints.
r.Route("/servers", func(r chi.Router) {
r.Get("/", api.listMCPServerConfigs)
r.Post("/", api.createMCPServerConfig)
r.Route("/{mcpServer}", func(r chi.Router) {
r.Get("/", api.getMCPServerConfig)
r.Patch("/", api.updateMCPServerConfig)
r.Delete("/", api.deleteMCPServerConfig)
// OAuth2 user flow
r.Get("/oauth2/connect", api.mcpServerOAuth2Connect)
r.Get("/oauth2/callback", api.mcpServerOAuth2Callback)
r.Delete("/oauth2/disconnect", api.mcpServerOAuth2Disconnect)
})
})
// This callback path is frozen because it is registered with OAuth2 providers.
r.Get("/servers/{mcpServer}/oauth2/callback", api.mcpServerOAuth2Callback)
// Disconnect stays outside organization routes so former organization
// members can delete their stored token after losing config read access.
r.Delete("/servers/{mcpServer}/oauth2/disconnect", api.mcpServerOAuth2Disconnect)
// MCP HTTP transport endpoint with mandatory authentication
r.Route("/http", func(r chi.Router) {
r.Use(httpmw.RequireExperimentWithDevBypass(api.Experiments, codersdk.ExperimentOAuth2, codersdk.ExperimentMCPServerHTTP))
+45 -31
View File
@@ -794,6 +794,7 @@ var (
rbac.ResourceChat.Type: {policy.ActionCreate, policy.ActionRead, policy.ActionUpdate, policy.ActionDelete},
rbac.ResourceWorkspace.Type: {policy.ActionRead, policy.ActionUpdate},
rbac.ResourceDeploymentConfig.Type: {policy.ActionRead},
rbac.ResourceMCPServerConfig.Type: {policy.ActionRead},
rbac.ResourceUser.Type: {policy.ActionReadPersonal},
}),
User: []rbac.Permission{},
@@ -2283,7 +2284,11 @@ func (q *querier) DeleteLicense(ctx context.Context, id int32) (int32, error) {
}
func (q *querier) DeleteMCPServerConfigByID(ctx context.Context, id uuid.UUID) error {
if err := q.authorizeContext(ctx, policy.ActionUpdate, rbac.ResourceDeploymentConfig); err != nil {
config, err := q.db.GetMCPServerConfigByID(ctx, id)
if err != nil {
return err
}
if err := q.authorizeContext(ctx, policy.ActionDelete, config); err != nil {
return err
}
return q.db.DeleteMCPServerConfigByID(ctx, id)
@@ -2296,6 +2301,17 @@ func (q *querier) DeleteMCPServerUserToken(ctx context.Context, arg database.Del
return q.db.DeleteMCPServerUserToken(ctx, arg)
}
func (q *querier) DeleteMCPServerUserTokensByConfigID(ctx context.Context, mcpServerConfigID uuid.UUID) error {
config, err := q.db.GetMCPServerConfigByID(ctx, mcpServerConfigID)
if err != nil {
return err
}
if err := q.authorizeContext(ctx, policy.ActionUpdate, config); err != nil {
return err
}
return q.db.DeleteMCPServerUserTokensByConfigID(ctx, mcpServerConfigID)
}
func (q *querier) DeleteOAuth2ProviderAppByClientID(ctx context.Context, id uuid.UUID) error {
if err := q.authorizeContext(ctx, policy.ActionDelete, rbac.ResourceOauth2App); err != nil {
return err
@@ -3763,11 +3779,12 @@ func (q *querier) GetEnabledChatModelConfigs(ctx context.Context) ([]database.Ge
return q.db.GetEnabledChatModelConfigs(ctx)
}
func (q *querier) GetEnabledMCPServerConfigs(ctx context.Context) ([]database.MCPServerConfig, error) {
if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceDeploymentConfig); err != nil {
return nil, err
}
return q.db.GetEnabledMCPServerConfigs(ctx)
func (q *querier) GetEnabledMCPServerConfigsByOrganization(ctx context.Context, organizationID uuid.UUID) ([]database.MCPServerConfig, error) {
return fetchWithPostFilter(q.auth, policy.ActionRead, q.db.GetEnabledMCPServerConfigsByOrganization)(ctx, organizationID)
}
func (q *querier) GetEnabledMCPServerConfigsByOrganizationAndIDs(ctx context.Context, arg database.GetEnabledMCPServerConfigsByOrganizationAndIDsParams) ([]database.MCPServerConfig, error) {
return fetchWithPostFilter(q.auth, policy.ActionRead, q.db.GetEnabledMCPServerConfigsByOrganizationAndIDs)(ctx, arg)
}
// GetExternalAgentTokensByTemplateID is used for scaletesting purposes; the
@@ -3842,11 +3859,8 @@ func (q *querier) GetFilteredInboxNotificationsByUserID(ctx context.Context, arg
return fetchWithPostFilter(q.auth, policy.ActionRead, q.db.GetFilteredInboxNotificationsByUserID)(ctx, arg)
}
func (q *querier) GetForcedMCPServerConfigs(ctx context.Context) ([]database.MCPServerConfig, error) {
if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceDeploymentConfig); err != nil {
return nil, err
}
return q.db.GetForcedMCPServerConfigs(ctx)
func (q *querier) GetForcedMCPServerConfigsByOrganization(ctx context.Context, organizationID uuid.UUID) ([]database.MCPServerConfig, error) {
return fetchWithPostFilter(q.auth, policy.ActionRead, q.db.GetForcedMCPServerConfigsByOrganization)(ctx, organizationID)
}
func (q *querier) GetGitSSHKey(ctx context.Context, userID uuid.UUID) (database.GitSSHKey, error) {
@@ -4048,31 +4062,23 @@ func (q *querier) GetLogoURL(ctx context.Context) (string, error) {
}
func (q *querier) GetMCPServerConfigByID(ctx context.Context, id uuid.UUID) (database.MCPServerConfig, error) {
if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceDeploymentConfig); err != nil {
return database.MCPServerConfig{}, err
}
return q.db.GetMCPServerConfigByID(ctx, id)
return fetch(q.log, q.auth, q.db.GetMCPServerConfigByID)(ctx, id)
}
func (q *querier) GetMCPServerConfigBySlug(ctx context.Context, slug string) (database.MCPServerConfig, error) {
if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceDeploymentConfig); err != nil {
return database.MCPServerConfig{}, err
}
return q.db.GetMCPServerConfigBySlug(ctx, slug)
func (q *querier) GetMCPServerConfigByIDForUpdate(ctx context.Context, id uuid.UUID) (database.MCPServerConfig, error) {
return fetch(q.log, q.auth, q.db.GetMCPServerConfigByIDForUpdate)(ctx, id)
}
func (q *querier) GetMCPServerConfigs(ctx context.Context) ([]database.MCPServerConfig, error) {
if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceDeploymentConfig); err != nil {
return nil, err
}
return q.db.GetMCPServerConfigs(ctx)
func (q *querier) GetMCPServerConfigByOrganizationAndSlug(ctx context.Context, arg database.GetMCPServerConfigByOrganizationAndSlugParams) (database.MCPServerConfig, error) {
return fetch(q.log, q.auth, q.db.GetMCPServerConfigByOrganizationAndSlug)(ctx, arg)
}
func (q *querier) GetMCPServerConfigsByIDs(ctx context.Context, ids []uuid.UUID) ([]database.MCPServerConfig, error) {
if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceDeploymentConfig); err != nil {
return nil, err
func (q *querier) GetMCPServerConfigsByOrganization(ctx context.Context, organizationID uuid.UUID) ([]database.MCPServerConfig, error) {
prepared, err := prepareSQLFilter(ctx, q.auth, policy.ActionRead, rbac.ResourceMCPServerConfig.Type)
if err != nil {
return nil, xerrors.Errorf("prepare sql filter: %w", err)
}
return q.db.GetMCPServerConfigsByIDs(ctx, ids)
return q.db.GetAuthorizedMCPServerConfigs(ctx, organizationID, prepared)
}
func (q *querier) GetMCPServerUserToken(ctx context.Context, arg database.GetMCPServerUserTokenParams) (database.MCPServerUserToken, error) {
@@ -6194,7 +6200,7 @@ func (q *querier) InsertLicense(ctx context.Context, arg database.InsertLicenseP
}
func (q *querier) InsertMCPServerConfig(ctx context.Context, arg database.InsertMCPServerConfigParams) (database.MCPServerConfig, error) {
if err := q.authorizeContext(ctx, policy.ActionUpdate, rbac.ResourceDeploymentConfig); err != nil {
if err := q.authorizeContext(ctx, policy.ActionCreate, rbac.ResourceMCPServerConfig.InOrg(arg.OrganizationID)); err != nil {
return database.MCPServerConfig{}, err
}
return q.db.InsertMCPServerConfig(ctx, arg)
@@ -7661,7 +7667,11 @@ func (q *querier) UpdateInboxNotificationReadStatus(ctx context.Context, args da
}
func (q *querier) UpdateMCPServerConfig(ctx context.Context, arg database.UpdateMCPServerConfigParams) (database.MCPServerConfig, error) {
if err := q.authorizeContext(ctx, policy.ActionUpdate, rbac.ResourceDeploymentConfig); err != nil {
config, err := q.db.GetMCPServerConfigByID(ctx, arg.ID)
if err != nil {
return database.MCPServerConfig{}, err
}
if err := q.authorizeContext(ctx, policy.ActionUpdate, config); err != nil {
return database.MCPServerConfig{}, err
}
return q.db.UpdateMCPServerConfig(ctx, arg)
@@ -9395,3 +9405,7 @@ func (q *querier) GetAuthorizedChats(ctx context.Context, arg database.GetChatsP
func (q *querier) GetAuthorizedChatsByChatFileID(ctx context.Context, fileID uuid.UUID, prepared rbac.PreparedAuthorized) ([]database.Chat, error) {
return q.db.GetAuthorizedChatsByChatFileID(ctx, fileID, prepared)
}
func (q *querier) GetAuthorizedMCPServerConfigs(ctx context.Context, organizationID uuid.UUID, prepared rbac.PreparedAuthorized) ([]database.MCPServerConfig, error) {
return q.db.GetAuthorizedMCPServerConfigs(ctx, organizationID, prepared)
}
+72 -36
View File
@@ -1632,10 +1632,11 @@ func (s *MethodTestSuite) TestChats() {
dbm.EXPECT().CleanupDeletedMCPServerIDsFromChats(gomock.Any()).Return(nil).AnyTimes()
check.Args().Asserts(rbac.ResourceChat, policy.ActionUpdate)
}))
s.Run("DeleteMCPServerConfigByID", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) {
id := uuid.New()
dbm.EXPECT().DeleteMCPServerConfigByID(gomock.Any(), id).Return(nil).AnyTimes()
check.Args(id).Asserts(rbac.ResourceDeploymentConfig, policy.ActionUpdate)
s.Run("DeleteMCPServerConfigByID", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
config := testutil.Fake(s.T(), faker, database.MCPServerConfig{})
dbm.EXPECT().GetMCPServerConfigByID(gomock.Any(), config.ID).Return(config, nil).AnyTimes()
dbm.EXPECT().DeleteMCPServerConfigByID(gomock.Any(), config.ID).Return(nil).AnyTimes()
check.Args(config.ID).Asserts(config, policy.ActionDelete)
}))
s.Run("DeleteMCPServerUserToken", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) {
arg := database.DeleteMCPServerUserTokenParams{
@@ -1645,41 +1646,68 @@ func (s *MethodTestSuite) TestChats() {
dbm.EXPECT().DeleteMCPServerUserToken(gomock.Any(), arg).Return(nil).AnyTimes()
check.Args(arg).Asserts(rbac.ResourceDeploymentConfig, policy.ActionUpdate)
}))
s.Run("GetEnabledMCPServerConfigs", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
configA := testutil.Fake(s.T(), faker, database.MCPServerConfig{})
configB := testutil.Fake(s.T(), faker, database.MCPServerConfig{})
dbm.EXPECT().GetEnabledMCPServerConfigs(gomock.Any()).Return([]database.MCPServerConfig{configA, configB}, nil).AnyTimes()
check.Args().Asserts(rbac.ResourceDeploymentConfig, policy.ActionRead).Returns([]database.MCPServerConfig{configA, configB})
s.Run("DeleteMCPServerUserTokensByConfigID", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
config := testutil.Fake(s.T(), faker, database.MCPServerConfig{})
dbm.EXPECT().GetMCPServerConfigByID(gomock.Any(), config.ID).Return(config, nil).AnyTimes()
dbm.EXPECT().DeleteMCPServerUserTokensByConfigID(gomock.Any(), config.ID).Return(nil).AnyTimes()
check.Args(config.ID).Asserts(config, policy.ActionUpdate)
}))
s.Run("GetForcedMCPServerConfigs", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
configA := testutil.Fake(s.T(), faker, database.MCPServerConfig{})
configB := testutil.Fake(s.T(), faker, database.MCPServerConfig{})
dbm.EXPECT().GetForcedMCPServerConfigs(gomock.Any()).Return([]database.MCPServerConfig{configA, configB}, nil).AnyTimes()
check.Args().Asserts(rbac.ResourceDeploymentConfig, policy.ActionRead).Returns([]database.MCPServerConfig{configA, configB})
s.Run("GetEnabledMCPServerConfigsByOrganization", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
orgID := uuid.New()
configA := testutil.Fake(s.T(), faker, database.MCPServerConfig{OrganizationID: orgID, Enabled: true})
configB := testutil.Fake(s.T(), faker, database.MCPServerConfig{OrganizationID: orgID, Enabled: true})
dbm.EXPECT().GetEnabledMCPServerConfigsByOrganization(gomock.Any(), orgID).Return([]database.MCPServerConfig{configA, configB}, nil).AnyTimes()
check.Args(orgID).Asserts(configA, policy.ActionRead, configB, policy.ActionRead).OutOfOrder().Returns([]database.MCPServerConfig{configA, configB})
}))
s.Run("GetForcedMCPServerConfigsByOrganization", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
orgID := uuid.New()
configA := testutil.Fake(s.T(), faker, database.MCPServerConfig{OrganizationID: orgID, Availability: "force_on"})
configB := testutil.Fake(s.T(), faker, database.MCPServerConfig{OrganizationID: orgID, Availability: "force_on"})
dbm.EXPECT().GetForcedMCPServerConfigsByOrganization(gomock.Any(), orgID).Return([]database.MCPServerConfig{configA, configB}, nil).AnyTimes()
check.Args(orgID).Asserts(configA, policy.ActionRead, configB, policy.ActionRead).OutOfOrder().Returns([]database.MCPServerConfig{configA, configB})
}))
s.Run("GetMCPServerConfigByID", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
config := testutil.Fake(s.T(), faker, database.MCPServerConfig{})
dbm.EXPECT().GetMCPServerConfigByID(gomock.Any(), config.ID).Return(config, nil).AnyTimes()
check.Args(config.ID).Asserts(rbac.ResourceDeploymentConfig, policy.ActionRead).Returns(config)
check.Args(config.ID).Asserts(config, policy.ActionRead).Returns(config)
}))
s.Run("GetMCPServerConfigBySlug", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
slug := "test-mcp-server"
config := testutil.Fake(s.T(), faker, database.MCPServerConfig{Slug: slug})
dbm.EXPECT().GetMCPServerConfigBySlug(gomock.Any(), slug).Return(config, nil).AnyTimes()
check.Args(slug).Asserts(rbac.ResourceDeploymentConfig, policy.ActionRead).Returns(config)
s.Run("GetMCPServerConfigByIDForUpdate", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
config := testutil.Fake(s.T(), faker, database.MCPServerConfig{})
dbm.EXPECT().GetMCPServerConfigByIDForUpdate(gomock.Any(), config.ID).Return(config, nil).AnyTimes()
check.Args(config.ID).Asserts(config, policy.ActionRead).Returns(config)
}))
s.Run("GetMCPServerConfigs", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
configA := testutil.Fake(s.T(), faker, database.MCPServerConfig{})
configB := testutil.Fake(s.T(), faker, database.MCPServerConfig{})
dbm.EXPECT().GetMCPServerConfigs(gomock.Any()).Return([]database.MCPServerConfig{configA, configB}, nil).AnyTimes()
check.Args().Asserts(rbac.ResourceDeploymentConfig, policy.ActionRead).Returns([]database.MCPServerConfig{configA, configB})
s.Run("GetMCPServerConfigByOrganizationAndSlug", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
arg := database.GetMCPServerConfigByOrganizationAndSlugParams{
OrganizationID: uuid.New(),
Slug: "test-mcp-server",
}
config := testutil.Fake(s.T(), faker, database.MCPServerConfig{OrganizationID: arg.OrganizationID, Slug: arg.Slug})
dbm.EXPECT().GetMCPServerConfigByOrganizationAndSlug(gomock.Any(), arg).Return(config, nil).AnyTimes()
check.Args(arg).Asserts(config, policy.ActionRead).Returns(config)
}))
s.Run("GetMCPServerConfigsByIDs", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
configA := testutil.Fake(s.T(), faker, database.MCPServerConfig{})
configB := testutil.Fake(s.T(), faker, database.MCPServerConfig{})
ids := []uuid.UUID{configA.ID, configB.ID}
dbm.EXPECT().GetMCPServerConfigsByIDs(gomock.Any(), ids).Return([]database.MCPServerConfig{configA, configB}, nil).AnyTimes()
check.Args(ids).Asserts(rbac.ResourceDeploymentConfig, policy.ActionRead).Returns([]database.MCPServerConfig{configA, configB})
s.Run("GetMCPServerConfigsByOrganization", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
orgID := uuid.New()
configA := testutil.Fake(s.T(), faker, database.MCPServerConfig{OrganizationID: orgID})
configB := testutil.Fake(s.T(), faker, database.MCPServerConfig{OrganizationID: orgID})
dbm.EXPECT().GetAuthorizedMCPServerConfigs(gomock.Any(), orgID, gomock.Any()).Return([]database.MCPServerConfig{configA, configB}, nil).AnyTimes()
check.Args(orgID).Asserts().Returns([]database.MCPServerConfig{configA, configB})
}))
s.Run("GetAuthorizedMCPServerConfigs", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
orgID := uuid.New()
configA := testutil.Fake(s.T(), faker, database.MCPServerConfig{OrganizationID: orgID})
configB := testutil.Fake(s.T(), faker, database.MCPServerConfig{OrganizationID: orgID})
dbm.EXPECT().GetAuthorizedMCPServerConfigs(gomock.Any(), orgID, gomock.Any()).Return([]database.MCPServerConfig{configA, configB}, nil).AnyTimes()
check.Args(orgID, emptyPreparedAuthorized{}).Asserts().Returns([]database.MCPServerConfig{configA, configB})
}))
s.Run("GetEnabledMCPServerConfigsByOrganizationAndIDs", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
arg := database.GetEnabledMCPServerConfigsByOrganizationAndIDsParams{
OrganizationID: uuid.New(),
IDs: []uuid.UUID{uuid.New(), uuid.New()},
}
configA := testutil.Fake(s.T(), faker, database.MCPServerConfig{ID: arg.IDs[0], OrganizationID: arg.OrganizationID})
configB := testutil.Fake(s.T(), faker, database.MCPServerConfig{ID: arg.IDs[1], OrganizationID: arg.OrganizationID})
dbm.EXPECT().GetEnabledMCPServerConfigsByOrganizationAndIDs(gomock.Any(), arg).Return([]database.MCPServerConfig{configA, configB}, nil).AnyTimes()
check.Args(arg).Asserts(configA, policy.ActionRead, configB, policy.ActionRead).OutOfOrder().Returns([]database.MCPServerConfig{configA, configB})
}))
s.Run("GetMCPServerUserToken", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
arg := database.GetMCPServerUserTokenParams{
@@ -1698,12 +1726,14 @@ func (s *MethodTestSuite) TestChats() {
}))
s.Run("InsertMCPServerConfig", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
arg := database.InsertMCPServerConfigParams{
DisplayName: "Test MCP Server",
Slug: "test-mcp-server",
ID: uuid.New(),
OrganizationID: uuid.New(),
DisplayName: "Test MCP Server",
Slug: "test-mcp-server",
}
config := testutil.Fake(s.T(), faker, database.MCPServerConfig{DisplayName: arg.DisplayName, Slug: arg.Slug})
config := testutil.Fake(s.T(), faker, database.MCPServerConfig{OrganizationID: arg.OrganizationID, DisplayName: arg.DisplayName, Slug: arg.Slug})
dbm.EXPECT().InsertMCPServerConfig(gomock.Any(), arg).Return(config, nil).AnyTimes()
check.Args(arg).Asserts(rbac.ResourceDeploymentConfig, policy.ActionUpdate).Returns(config)
check.Args(arg).Asserts(rbac.ResourceMCPServerConfig.InOrg(arg.OrganizationID), policy.ActionCreate).Returns(config)
}))
s.Run("UpdateChatMCPServerIDs", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
chat := testutil.Fake(s.T(), faker, database.Chat{})
@@ -1754,8 +1784,9 @@ func (s *MethodTestSuite) TestChats() {
DisplayName: "Updated MCP Server",
Slug: "updated-mcp-server",
}
dbm.EXPECT().GetMCPServerConfigByID(gomock.Any(), config.ID).Return(config, nil).AnyTimes()
dbm.EXPECT().UpdateMCPServerConfig(gomock.Any(), arg).Return(config, nil).AnyTimes()
check.Args(arg).Asserts(rbac.ResourceDeploymentConfig, policy.ActionUpdate).Returns(config)
check.Args(arg).Asserts(config, 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{})
@@ -7587,6 +7618,11 @@ func TestAsChatd(t *testing.T) {
// Cannot access provisioner daemons.
err = auth.Authorize(ctx, actor, policy.ActionRead, rbac.ResourceProvisionerDaemon)
require.Error(t, err, "provisioner daemon read should be denied")
// Cannot access organizations; MCP server config resolution is
// strictly org-scoped and needs no organization reads.
err = auth.Authorize(ctx, actor, policy.ActionRead, rbac.ResourceOrganization)
require.Error(t, err, "organization read should be denied")
})
}
+11
View File
@@ -328,6 +328,15 @@ func ChatProvider(t testing.TB, db database.Store, seed database.ChatProvider, m
func MCPServerConfig(t testing.TB, db database.Store, seed database.MCPServerConfig) database.MCPServerConfig {
t.Helper()
// New configs belong to the default organization, matching the
// org-less shape they had before configs became org-scoped.
organizationID := seed.OrganizationID
if organizationID == uuid.Nil {
defaultOrg, err := db.GetDefaultOrganization(genCtx)
require.NoError(t, err, "get default organization")
organizationID = defaultOrg.ID
}
// CreatedBy and UpdatedBy are user FKs, so default fixtures create a user.
createdBy := seed.CreatedBy.UUID
if createdBy == uuid.Nil {
@@ -339,6 +348,8 @@ func MCPServerConfig(t testing.TB, db database.Store, seed database.MCPServerCon
}
cfg, err := db.InsertMCPServerConfig(genCtx, database.InsertMCPServerConfigParams{
ID: takeFirst(seed.ID, uuid.New()),
OrganizationID: organizationID,
DisplayName: takeFirst(seed.DisplayName, "Test MCP Server"),
Slug: takeFirst(seed.Slug, testutil.GetRandomName(t)),
Description: seed.Description,
+44 -20
View File
@@ -625,6 +625,14 @@ func (m queryMetricsStore) DeleteMCPServerUserToken(ctx context.Context, arg dat
return r0
}
func (m queryMetricsStore) DeleteMCPServerUserTokensByConfigID(ctx context.Context, mcpServerConfigID uuid.UUID) error {
start := time.Now()
r0 := m.s.DeleteMCPServerUserTokensByConfigID(ctx, mcpServerConfigID)
m.queryLatencies.WithLabelValues("DeleteMCPServerUserTokensByConfigID").Observe(time.Since(start).Seconds())
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "DeleteMCPServerUserTokensByConfigID").Inc()
return r0
}
func (m queryMetricsStore) DeleteOAuth2ProviderAppByClientID(ctx context.Context, id uuid.UUID) error {
start := time.Now()
r0 := m.s.DeleteOAuth2ProviderAppByClientID(ctx, id)
@@ -2009,11 +2017,19 @@ func (m queryMetricsStore) GetEnabledChatModelConfigs(ctx context.Context) ([]da
return r0, r1
}
func (m queryMetricsStore) GetEnabledMCPServerConfigs(ctx context.Context) ([]database.MCPServerConfig, error) {
func (m queryMetricsStore) GetEnabledMCPServerConfigsByOrganization(ctx context.Context, organizationID uuid.UUID) ([]database.MCPServerConfig, error) {
start := time.Now()
r0, r1 := m.s.GetEnabledMCPServerConfigs(ctx)
m.queryLatencies.WithLabelValues("GetEnabledMCPServerConfigs").Observe(time.Since(start).Seconds())
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetEnabledMCPServerConfigs").Inc()
r0, r1 := m.s.GetEnabledMCPServerConfigsByOrganization(ctx, organizationID)
m.queryLatencies.WithLabelValues("GetEnabledMCPServerConfigsByOrganization").Observe(time.Since(start).Seconds())
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetEnabledMCPServerConfigsByOrganization").Inc()
return r0, r1
}
func (m queryMetricsStore) GetEnabledMCPServerConfigsByOrganizationAndIDs(ctx context.Context, arg database.GetEnabledMCPServerConfigsByOrganizationAndIDsParams) ([]database.MCPServerConfig, error) {
start := time.Now()
r0, r1 := m.s.GetEnabledMCPServerConfigsByOrganizationAndIDs(ctx, arg)
m.queryLatencies.WithLabelValues("GetEnabledMCPServerConfigsByOrganizationAndIDs").Observe(time.Since(start).Seconds())
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetEnabledMCPServerConfigsByOrganizationAndIDs").Inc()
return r0, r1
}
@@ -2081,11 +2097,11 @@ func (m queryMetricsStore) GetFilteredInboxNotificationsByUserID(ctx context.Con
return r0, r1
}
func (m queryMetricsStore) GetForcedMCPServerConfigs(ctx context.Context) ([]database.MCPServerConfig, error) {
func (m queryMetricsStore) GetForcedMCPServerConfigsByOrganization(ctx context.Context, organizationID uuid.UUID) ([]database.MCPServerConfig, error) {
start := time.Now()
r0, r1 := m.s.GetForcedMCPServerConfigs(ctx)
m.queryLatencies.WithLabelValues("GetForcedMCPServerConfigs").Observe(time.Since(start).Seconds())
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetForcedMCPServerConfigs").Inc()
r0, r1 := m.s.GetForcedMCPServerConfigsByOrganization(ctx, organizationID)
m.queryLatencies.WithLabelValues("GetForcedMCPServerConfigsByOrganization").Observe(time.Since(start).Seconds())
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetForcedMCPServerConfigsByOrganization").Inc()
return r0, r1
}
@@ -2321,27 +2337,27 @@ func (m queryMetricsStore) GetMCPServerConfigByID(ctx context.Context, id uuid.U
return r0, r1
}
func (m queryMetricsStore) GetMCPServerConfigBySlug(ctx context.Context, slug string) (database.MCPServerConfig, error) {
func (m queryMetricsStore) GetMCPServerConfigByIDForUpdate(ctx context.Context, id uuid.UUID) (database.MCPServerConfig, error) {
start := time.Now()
r0, r1 := m.s.GetMCPServerConfigBySlug(ctx, slug)
m.queryLatencies.WithLabelValues("GetMCPServerConfigBySlug").Observe(time.Since(start).Seconds())
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetMCPServerConfigBySlug").Inc()
r0, r1 := m.s.GetMCPServerConfigByIDForUpdate(ctx, id)
m.queryLatencies.WithLabelValues("GetMCPServerConfigByIDForUpdate").Observe(time.Since(start).Seconds())
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetMCPServerConfigByIDForUpdate").Inc()
return r0, r1
}
func (m queryMetricsStore) GetMCPServerConfigs(ctx context.Context) ([]database.MCPServerConfig, error) {
func (m queryMetricsStore) GetMCPServerConfigByOrganizationAndSlug(ctx context.Context, arg database.GetMCPServerConfigByOrganizationAndSlugParams) (database.MCPServerConfig, error) {
start := time.Now()
r0, r1 := m.s.GetMCPServerConfigs(ctx)
m.queryLatencies.WithLabelValues("GetMCPServerConfigs").Observe(time.Since(start).Seconds())
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetMCPServerConfigs").Inc()
r0, r1 := m.s.GetMCPServerConfigByOrganizationAndSlug(ctx, arg)
m.queryLatencies.WithLabelValues("GetMCPServerConfigByOrganizationAndSlug").Observe(time.Since(start).Seconds())
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetMCPServerConfigByOrganizationAndSlug").Inc()
return r0, r1
}
func (m queryMetricsStore) GetMCPServerConfigsByIDs(ctx context.Context, ids []uuid.UUID) ([]database.MCPServerConfig, error) {
func (m queryMetricsStore) GetMCPServerConfigsByOrganization(ctx context.Context, organizationID uuid.UUID) ([]database.MCPServerConfig, error) {
start := time.Now()
r0, r1 := m.s.GetMCPServerConfigsByIDs(ctx, ids)
m.queryLatencies.WithLabelValues("GetMCPServerConfigsByIDs").Observe(time.Since(start).Seconds())
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetMCPServerConfigsByIDs").Inc()
r0, r1 := m.s.GetMCPServerConfigsByOrganization(ctx, organizationID)
m.queryLatencies.WithLabelValues("GetMCPServerConfigsByOrganization").Observe(time.Since(start).Seconds())
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetMCPServerConfigsByOrganization").Inc()
return r0, r1
}
@@ -6792,3 +6808,11 @@ func (m queryMetricsStore) GetAuthorizedChatsByChatFileID(ctx context.Context, f
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetAuthorizedChatsByChatFileID").Inc()
return r0, r1
}
func (m queryMetricsStore) GetAuthorizedMCPServerConfigs(ctx context.Context, organizationID uuid.UUID, prepared rbac.PreparedAuthorized) ([]database.MCPServerConfig, error) {
start := time.Now()
r0, r1 := m.s.GetAuthorizedMCPServerConfigs(ctx, organizationID, prepared)
m.queryLatencies.WithLabelValues("GetAuthorizedMCPServerConfigs").Observe(time.Since(start).Seconds())
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetAuthorizedMCPServerConfigs").Inc()
return r0, r1
}
+83 -39
View File
@@ -1035,6 +1035,20 @@ func (mr *MockStoreMockRecorder) DeleteMCPServerUserToken(ctx, arg any) *gomock.
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteMCPServerUserToken", reflect.TypeOf((*MockStore)(nil).DeleteMCPServerUserToken), ctx, arg)
}
// DeleteMCPServerUserTokensByConfigID mocks base method.
func (m *MockStore) DeleteMCPServerUserTokensByConfigID(ctx context.Context, mcpServerConfigID uuid.UUID) error {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "DeleteMCPServerUserTokensByConfigID", ctx, mcpServerConfigID)
ret0, _ := ret[0].(error)
return ret0
}
// DeleteMCPServerUserTokensByConfigID indicates an expected call of DeleteMCPServerUserTokensByConfigID.
func (mr *MockStoreMockRecorder) DeleteMCPServerUserTokensByConfigID(ctx, mcpServerConfigID any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteMCPServerUserTokensByConfigID", reflect.TypeOf((*MockStore)(nil).DeleteMCPServerUserTokensByConfigID), ctx, mcpServerConfigID)
}
// DeleteOAuth2ProviderAppByClientID mocks base method.
func (m *MockStore) DeleteOAuth2ProviderAppByClientID(ctx context.Context, id uuid.UUID) error {
m.ctrl.T.Helper()
@@ -2505,6 +2519,21 @@ func (mr *MockStoreMockRecorder) GetAuthorizedConnectionLogsOffset(ctx, arg, pre
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAuthorizedConnectionLogsOffset", reflect.TypeOf((*MockStore)(nil).GetAuthorizedConnectionLogsOffset), ctx, arg, prepared)
}
// GetAuthorizedMCPServerConfigs mocks base method.
func (m *MockStore) GetAuthorizedMCPServerConfigs(ctx context.Context, organizationID uuid.UUID, prepared rbac.PreparedAuthorized) ([]database.MCPServerConfig, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "GetAuthorizedMCPServerConfigs", ctx, organizationID, prepared)
ret0, _ := ret[0].([]database.MCPServerConfig)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// GetAuthorizedMCPServerConfigs indicates an expected call of GetAuthorizedMCPServerConfigs.
func (mr *MockStoreMockRecorder) GetAuthorizedMCPServerConfigs(ctx, organizationID, prepared any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAuthorizedMCPServerConfigs", reflect.TypeOf((*MockStore)(nil).GetAuthorizedMCPServerConfigs), ctx, organizationID, prepared)
}
// GetAuthorizedTemplates mocks base method.
func (m *MockStore) GetAuthorizedTemplates(ctx context.Context, arg database.GetTemplatesWithFilterParams, prepared rbac.PreparedAuthorized) ([]database.Template, error) {
m.ctrl.T.Helper()
@@ -3735,19 +3764,34 @@ func (mr *MockStoreMockRecorder) GetEnabledChatModelConfigs(ctx any) *gomock.Cal
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetEnabledChatModelConfigs", reflect.TypeOf((*MockStore)(nil).GetEnabledChatModelConfigs), ctx)
}
// GetEnabledMCPServerConfigs mocks base method.
func (m *MockStore) GetEnabledMCPServerConfigs(ctx context.Context) ([]database.MCPServerConfig, error) {
// GetEnabledMCPServerConfigsByOrganization mocks base method.
func (m *MockStore) GetEnabledMCPServerConfigsByOrganization(ctx context.Context, organizationID uuid.UUID) ([]database.MCPServerConfig, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "GetEnabledMCPServerConfigs", ctx)
ret := m.ctrl.Call(m, "GetEnabledMCPServerConfigsByOrganization", ctx, organizationID)
ret0, _ := ret[0].([]database.MCPServerConfig)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// GetEnabledMCPServerConfigs indicates an expected call of GetEnabledMCPServerConfigs.
func (mr *MockStoreMockRecorder) GetEnabledMCPServerConfigs(ctx any) *gomock.Call {
// GetEnabledMCPServerConfigsByOrganization indicates an expected call of GetEnabledMCPServerConfigsByOrganization.
func (mr *MockStoreMockRecorder) GetEnabledMCPServerConfigsByOrganization(ctx, organizationID any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetEnabledMCPServerConfigs", reflect.TypeOf((*MockStore)(nil).GetEnabledMCPServerConfigs), ctx)
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetEnabledMCPServerConfigsByOrganization", reflect.TypeOf((*MockStore)(nil).GetEnabledMCPServerConfigsByOrganization), ctx, organizationID)
}
// GetEnabledMCPServerConfigsByOrganizationAndIDs mocks base method.
func (m *MockStore) GetEnabledMCPServerConfigsByOrganizationAndIDs(ctx context.Context, arg database.GetEnabledMCPServerConfigsByOrganizationAndIDsParams) ([]database.MCPServerConfig, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "GetEnabledMCPServerConfigsByOrganizationAndIDs", ctx, arg)
ret0, _ := ret[0].([]database.MCPServerConfig)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// GetEnabledMCPServerConfigsByOrganizationAndIDs indicates an expected call of GetEnabledMCPServerConfigsByOrganizationAndIDs.
func (mr *MockStoreMockRecorder) GetEnabledMCPServerConfigsByOrganizationAndIDs(ctx, arg any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetEnabledMCPServerConfigsByOrganizationAndIDs", reflect.TypeOf((*MockStore)(nil).GetEnabledMCPServerConfigsByOrganizationAndIDs), ctx, arg)
}
// GetExternalAgentTokensByTemplateID mocks base method.
@@ -3870,19 +3914,19 @@ func (mr *MockStoreMockRecorder) GetFilteredInboxNotificationsByUserID(ctx, arg
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetFilteredInboxNotificationsByUserID", reflect.TypeOf((*MockStore)(nil).GetFilteredInboxNotificationsByUserID), ctx, arg)
}
// GetForcedMCPServerConfigs mocks base method.
func (m *MockStore) GetForcedMCPServerConfigs(ctx context.Context) ([]database.MCPServerConfig, error) {
// GetForcedMCPServerConfigsByOrganization mocks base method.
func (m *MockStore) GetForcedMCPServerConfigsByOrganization(ctx context.Context, organizationID uuid.UUID) ([]database.MCPServerConfig, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "GetForcedMCPServerConfigs", ctx)
ret := m.ctrl.Call(m, "GetForcedMCPServerConfigsByOrganization", ctx, organizationID)
ret0, _ := ret[0].([]database.MCPServerConfig)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// GetForcedMCPServerConfigs indicates an expected call of GetForcedMCPServerConfigs.
func (mr *MockStoreMockRecorder) GetForcedMCPServerConfigs(ctx any) *gomock.Call {
// GetForcedMCPServerConfigsByOrganization indicates an expected call of GetForcedMCPServerConfigsByOrganization.
func (mr *MockStoreMockRecorder) GetForcedMCPServerConfigsByOrganization(ctx, organizationID any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetForcedMCPServerConfigs", reflect.TypeOf((*MockStore)(nil).GetForcedMCPServerConfigs), ctx)
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetForcedMCPServerConfigsByOrganization", reflect.TypeOf((*MockStore)(nil).GetForcedMCPServerConfigsByOrganization), ctx, organizationID)
}
// GetGitSSHKey mocks base method.
@@ -4320,49 +4364,49 @@ func (mr *MockStoreMockRecorder) GetMCPServerConfigByID(ctx, id any) *gomock.Cal
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetMCPServerConfigByID", reflect.TypeOf((*MockStore)(nil).GetMCPServerConfigByID), ctx, id)
}
// GetMCPServerConfigBySlug mocks base method.
func (m *MockStore) GetMCPServerConfigBySlug(ctx context.Context, slug string) (database.MCPServerConfig, error) {
// GetMCPServerConfigByIDForUpdate mocks base method.
func (m *MockStore) GetMCPServerConfigByIDForUpdate(ctx context.Context, id uuid.UUID) (database.MCPServerConfig, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "GetMCPServerConfigBySlug", ctx, slug)
ret := m.ctrl.Call(m, "GetMCPServerConfigByIDForUpdate", ctx, id)
ret0, _ := ret[0].(database.MCPServerConfig)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// GetMCPServerConfigBySlug indicates an expected call of GetMCPServerConfigBySlug.
func (mr *MockStoreMockRecorder) GetMCPServerConfigBySlug(ctx, slug any) *gomock.Call {
// GetMCPServerConfigByIDForUpdate indicates an expected call of GetMCPServerConfigByIDForUpdate.
func (mr *MockStoreMockRecorder) GetMCPServerConfigByIDForUpdate(ctx, id any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetMCPServerConfigBySlug", reflect.TypeOf((*MockStore)(nil).GetMCPServerConfigBySlug), ctx, slug)
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetMCPServerConfigByIDForUpdate", reflect.TypeOf((*MockStore)(nil).GetMCPServerConfigByIDForUpdate), ctx, id)
}
// GetMCPServerConfigs mocks base method.
func (m *MockStore) GetMCPServerConfigs(ctx context.Context) ([]database.MCPServerConfig, error) {
// GetMCPServerConfigByOrganizationAndSlug mocks base method.
func (m *MockStore) GetMCPServerConfigByOrganizationAndSlug(ctx context.Context, arg database.GetMCPServerConfigByOrganizationAndSlugParams) (database.MCPServerConfig, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "GetMCPServerConfigs", ctx)
ret := m.ctrl.Call(m, "GetMCPServerConfigByOrganizationAndSlug", ctx, arg)
ret0, _ := ret[0].(database.MCPServerConfig)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// GetMCPServerConfigByOrganizationAndSlug indicates an expected call of GetMCPServerConfigByOrganizationAndSlug.
func (mr *MockStoreMockRecorder) GetMCPServerConfigByOrganizationAndSlug(ctx, arg any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetMCPServerConfigByOrganizationAndSlug", reflect.TypeOf((*MockStore)(nil).GetMCPServerConfigByOrganizationAndSlug), ctx, arg)
}
// GetMCPServerConfigsByOrganization mocks base method.
func (m *MockStore) GetMCPServerConfigsByOrganization(ctx context.Context, organizationID uuid.UUID) ([]database.MCPServerConfig, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "GetMCPServerConfigsByOrganization", ctx, organizationID)
ret0, _ := ret[0].([]database.MCPServerConfig)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// GetMCPServerConfigs indicates an expected call of GetMCPServerConfigs.
func (mr *MockStoreMockRecorder) GetMCPServerConfigs(ctx any) *gomock.Call {
// GetMCPServerConfigsByOrganization indicates an expected call of GetMCPServerConfigsByOrganization.
func (mr *MockStoreMockRecorder) GetMCPServerConfigsByOrganization(ctx, organizationID any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetMCPServerConfigs", reflect.TypeOf((*MockStore)(nil).GetMCPServerConfigs), ctx)
}
// GetMCPServerConfigsByIDs mocks base method.
func (m *MockStore) GetMCPServerConfigsByIDs(ctx context.Context, ids []uuid.UUID) ([]database.MCPServerConfig, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "GetMCPServerConfigsByIDs", ctx, ids)
ret0, _ := ret[0].([]database.MCPServerConfig)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// GetMCPServerConfigsByIDs indicates an expected call of GetMCPServerConfigsByIDs.
func (mr *MockStoreMockRecorder) GetMCPServerConfigsByIDs(ctx, ids any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetMCPServerConfigsByIDs", reflect.TypeOf((*MockStore)(nil).GetMCPServerConfigsByIDs), ctx, ids)
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetMCPServerConfigsByOrganization", reflect.TypeOf((*MockStore)(nil).GetMCPServerConfigsByOrganization), ctx, organizationID)
}
// GetMCPServerUserToken mocks base method.
+14 -3
View File
@@ -273,7 +273,12 @@ CREATE TYPE api_key_scope AS ENUM (
'workspace_build_orchestration:create',
'workspace_build_orchestration:delete',
'workspace_build_orchestration:read',
'workspace_build_orchestration:update'
'workspace_build_orchestration:update',
'mcp_server_config:*',
'mcp_server_config:create',
'mcp_server_config:read',
'mcp_server_config:update',
'mcp_server_config:delete'
);
CREATE TYPE app_sharing_level AS ENUM (
@@ -2514,6 +2519,7 @@ CREATE TABLE mcp_server_configs (
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,
organization_id uuid 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])))
@@ -4404,10 +4410,10 @@ ALTER TABLE ONLY licenses
ADD CONSTRAINT licenses_pkey PRIMARY KEY (id);
ALTER TABLE ONLY mcp_server_configs
ADD CONSTRAINT mcp_server_configs_pkey PRIMARY KEY (id);
ADD CONSTRAINT mcp_server_configs_organization_id_slug_key UNIQUE (organization_id, slug);
ALTER TABLE ONLY mcp_server_configs
ADD CONSTRAINT mcp_server_configs_slug_key UNIQUE (slug);
ADD CONSTRAINT mcp_server_configs_pkey PRIMARY KEY (id);
ALTER TABLE ONLY mcp_server_user_tokens
ADD CONSTRAINT mcp_server_user_tokens_mcp_server_config_id_user_id_key UNIQUE (mcp_server_config_id, user_id);
@@ -4872,6 +4878,8 @@ CREATE INDEX idx_mcp_server_configs_enabled ON mcp_server_configs USING btree (e
CREATE INDEX idx_mcp_server_configs_forced ON mcp_server_configs USING btree (enabled, availability) WHERE ((enabled = true) AND (availability = 'force_on'::text));
CREATE INDEX idx_mcp_server_configs_organization_id ON mcp_server_configs USING btree (organization_id);
CREATE INDEX idx_mcp_server_user_tokens_user_id ON mcp_server_user_tokens USING btree (user_id);
CREATE INDEX idx_notification_messages_status ON notification_messages USING btree (status);
@@ -5296,6 +5304,9 @@ ALTER TABLE ONLY mcp_server_configs
ALTER TABLE ONLY mcp_server_configs
ADD CONSTRAINT mcp_server_configs_oauth2_client_secret_key_id_fkey FOREIGN KEY (oauth2_client_secret_key_id) REFERENCES dbcrypt_keys(active_key_digest);
ALTER TABLE ONLY mcp_server_configs
ADD CONSTRAINT mcp_server_configs_organization_id_fkey FOREIGN KEY (organization_id) REFERENCES organizations(id) ON DELETE CASCADE;
ALTER TABLE ONLY mcp_server_configs
ADD CONSTRAINT mcp_server_configs_updated_by_fkey FOREIGN KEY (updated_by) REFERENCES users(id) ON DELETE SET NULL;
+1
View File
@@ -60,6 +60,7 @@ const (
ForeignKeyMcpServerConfigsCreatedBy ForeignKeyConstraint = "mcp_server_configs_created_by_fkey" // ALTER TABLE ONLY mcp_server_configs ADD CONSTRAINT mcp_server_configs_created_by_fkey FOREIGN KEY (created_by) REFERENCES users(id) ON DELETE SET NULL;
ForeignKeyMcpServerConfigsCustomHeadersKeyID ForeignKeyConstraint = "mcp_server_configs_custom_headers_key_id_fkey" // ALTER TABLE ONLY mcp_server_configs ADD CONSTRAINT mcp_server_configs_custom_headers_key_id_fkey FOREIGN KEY (custom_headers_key_id) REFERENCES dbcrypt_keys(active_key_digest);
ForeignKeyMcpServerConfigsOauth2ClientSecretKeyID ForeignKeyConstraint = "mcp_server_configs_oauth2_client_secret_key_id_fkey" // ALTER TABLE ONLY mcp_server_configs ADD CONSTRAINT mcp_server_configs_oauth2_client_secret_key_id_fkey FOREIGN KEY (oauth2_client_secret_key_id) REFERENCES dbcrypt_keys(active_key_digest);
ForeignKeyMcpServerConfigsOrganizationID ForeignKeyConstraint = "mcp_server_configs_organization_id_fkey" // ALTER TABLE ONLY mcp_server_configs ADD CONSTRAINT mcp_server_configs_organization_id_fkey FOREIGN KEY (organization_id) REFERENCES organizations(id) ON DELETE CASCADE;
ForeignKeyMcpServerConfigsUpdatedBy ForeignKeyConstraint = "mcp_server_configs_updated_by_fkey" // ALTER TABLE ONLY mcp_server_configs ADD CONSTRAINT mcp_server_configs_updated_by_fkey FOREIGN KEY (updated_by) REFERENCES users(id) ON DELETE SET NULL;
ForeignKeyMcpServerUserTokensAccessTokenKeyID ForeignKeyConstraint = "mcp_server_user_tokens_access_token_key_id_fkey" // ALTER TABLE ONLY mcp_server_user_tokens ADD CONSTRAINT mcp_server_user_tokens_access_token_key_id_fkey FOREIGN KEY (access_token_key_id) REFERENCES dbcrypt_keys(active_key_digest);
ForeignKeyMcpServerUserTokensMcpServerConfigID ForeignKeyConstraint = "mcp_server_user_tokens_mcp_server_config_id_fkey" // ALTER TABLE ONLY mcp_server_user_tokens ADD CONSTRAINT mcp_server_user_tokens_mcp_server_config_id_fkey FOREIGN KEY (mcp_server_config_id) REFERENCES mcp_server_configs(id) ON DELETE CASCADE;
@@ -0,0 +1,16 @@
-- Configs created outside the default organization cannot move to the
-- deployment-wide table because slugs may collide across organizations.
-- Delete them; the delete trigger from 000510 removes their IDs from chats.
DELETE FROM mcp_server_configs
WHERE organization_id != (
SELECT id FROM organizations WHERE is_default = true LIMIT 1
);
DROP INDEX idx_mcp_server_configs_organization_id;
ALTER TABLE mcp_server_configs
DROP CONSTRAINT mcp_server_configs_organization_id_slug_key,
DROP COLUMN organization_id,
ADD CONSTRAINT mcp_server_configs_slug_key UNIQUE (slug);
-- Enum values cannot be removed safely from api_key_scope.
@@ -0,0 +1,21 @@
ALTER TYPE api_key_scope ADD VALUE IF NOT EXISTS 'mcp_server_config:*';
ALTER TYPE api_key_scope ADD VALUE IF NOT EXISTS 'mcp_server_config:create';
ALTER TYPE api_key_scope ADD VALUE IF NOT EXISTS 'mcp_server_config:read';
ALTER TYPE api_key_scope ADD VALUE IF NOT EXISTS 'mcp_server_config:update';
ALTER TYPE api_key_scope ADD VALUE IF NOT EXISTS 'mcp_server_config:delete';
ALTER TABLE mcp_server_configs
ADD COLUMN organization_id UUID REFERENCES organizations(id) ON DELETE CASCADE;
-- The deployment-wide originals become the default organization's servers,
-- credentials intact. Other organizations start with no MCP servers.
UPDATE mcp_server_configs
SET organization_id = (SELECT id FROM organizations WHERE is_default = true LIMIT 1);
ALTER TABLE mcp_server_configs
ALTER COLUMN organization_id SET NOT NULL,
DROP CONSTRAINT mcp_server_configs_slug_key,
ADD CONSTRAINT mcp_server_configs_organization_id_slug_key UNIQUE (organization_id, slug);
CREATE INDEX idx_mcp_server_configs_organization_id
ON mcp_server_configs (organization_id);
+247
View File
@@ -117,6 +117,21 @@ func testSQLDB(t testing.TB) *sql.DB {
return db
}
func stepMigrationsUpTo(t *testing.T, next func() (version uint, more bool, err error), target uint) {
t.Helper()
for {
version, more, err := next()
require.NoError(t, err)
if !more {
t.Fatalf("migration %d not found", target)
}
if version == target {
return
}
}
}
// paralleltest linter doesn't correctly handle table-driven tests (https://github.com/kunwardeep/paralleltest/issues/8)
// nolint:paralleltest
func TestCheckLatestVersion(t *testing.T) {
@@ -2992,3 +3007,235 @@ func TestMigration000566OAuth2AuthMethodBackfill(t *testing.T) {
require.Equal(t, "confidential", stillConfidential,
"the backfill aligns the declaration to what is enforced, so the enforced value must be unchanged")
}
func TestMigration000574MCPServerConfigsOrganizationID(t *testing.T) {
t.Parallel()
const priorMigrationVersion = 573
sqlDB := testSQLDB(t)
next, err := migrations.Stepper(sqlDB)
require.NoError(t, err)
stepMigrationsUpTo(t, next, priorMigrationVersion)
ctx := testutil.Context(t, testutil.WaitSuperLong)
now := time.Now().UTC().Truncate(time.Microsecond)
var defaultOrgID uuid.UUID
err = sqlDB.QueryRowContext(ctx, `SELECT id FROM organizations WHERE is_default = true`).Scan(&defaultOrgID)
require.NoError(t, err)
otherOrgID := uuid.New()
_, err = sqlDB.ExecContext(ctx, `
INSERT INTO organizations (
id, name, display_name, description, icon, created_at, updated_at,
is_default, deleted, default_org_member_roles
) VALUES ($1, 'migration-568-org', 'Migration 568 Org', '', '', $2, $2, false, false, '{}')
`, otherOrgID, now)
require.NoError(t, err)
userID := uuid.New()
_, err = sqlDB.ExecContext(ctx, `
INSERT INTO users (
id, username, email, hashed_password, created_at, updated_at,
status, rbac_roles, login_type
) VALUES ($1, 'migration-568-user', 'migration-568@example.com', ''::bytea, $2, $2, 'active', '{}', 'password')
`, userID, now)
require.NoError(t, err)
const keyDigest = "migration-568-key"
_, err = sqlDB.ExecContext(ctx, `
INSERT INTO dbcrypt_keys (number, active_key_digest, test)
VALUES (568000, $1, 'migration-568-test')
`, keyDigest)
require.NoError(t, err)
providerID := uuid.New()
modelConfigID := uuid.New()
_, err = sqlDB.ExecContext(ctx, `
INSERT INTO ai_providers (
id, type, name, display_name, enabled, base_url, created_at, updated_at
) VALUES ($1, 'openai', 'migration-568-provider', 'Migration 568 Provider', true, 'https://provider.example.com', $2, $2)
`, providerID, now)
require.NoError(t, err)
_, err = sqlDB.ExecContext(ctx, `
INSERT INTO chat_model_configs (
id, model, display_name, ai_provider_id, context_limit,
compression_threshold, created_at, updated_at
) VALUES ($1, 'migration-568-model', 'Migration 568 Model', $2, 128000, 70, $3, $3)
`, modelConfigID, providerID, now)
require.NoError(t, err)
type configSeed struct {
id uuid.UUID
slug string
authType string
oauth2ClientID string
oauth2ClientSecret string
oauth2ClientSecretKeyID sql.NullString
oauth2AuthURL string
oauth2TokenURL string
oauth2RevocationURL string
oauth2Scopes string
apiKeyHeader string
apiKeyValue string
apiKeyValueKeyID sql.NullString
customHeaders string
customHeadersKeyID sql.NullString
}
keyID := sql.NullString{String: keyDigest, Valid: true}
configs := []configSeed{
{id: uuid.New(), slug: "migration-568-none", authType: "none", apiKeyHeader: "Authorization", customHeaders: "{}"},
// Every credential column carries ciphertext to prove the backfill
// leaves rows byte-identical apart from organization_id.
{
id: uuid.New(),
slug: "migration-568-oauth2",
authType: "oauth2",
oauth2ClientID: "oauth-client-id",
oauth2ClientSecret: "oauth-secret-ciphertext",
oauth2ClientSecretKeyID: keyID,
oauth2AuthURL: "https://oauth.example.com/authorize",
oauth2TokenURL: "https://oauth.example.com/token",
oauth2RevocationURL: "https://oauth.example.com/revoke",
oauth2Scopes: "openid profile",
apiKeyHeader: "X-API-Key",
apiKeyValue: "api-key-ciphertext",
apiKeyValueKeyID: keyID,
customHeaders: "custom-headers-ciphertext",
customHeadersKeyID: keyID,
},
}
originalJSON := make(map[uuid.UUID]string, len(configs))
for _, config := range configs {
_, err = sqlDB.ExecContext(ctx, `
INSERT INTO mcp_server_configs (
id, display_name, slug, description, url, auth_type,
oauth2_client_id, oauth2_client_secret, oauth2_client_secret_key_id,
oauth2_auth_url, oauth2_token_url, oauth2_revocation_url, oauth2_scopes,
api_key_header, api_key_value, api_key_value_key_id,
custom_headers, custom_headers_key_id, availability, enabled,
created_by, updated_by, created_at, updated_at
) VALUES (
$1, $2, $3, $4, $5, $6,
$7, $8, $9, $10, $11, $12, $13,
$14, $15, $16, $17, $18, 'default_on', true,
$19, $19, $20, $20
)
`,
config.id, "Migration 568 "+config.authType, config.slug, "migration 568 config", "https://mcp.example.com/"+config.slug, config.authType,
config.oauth2ClientID, config.oauth2ClientSecret, config.oauth2ClientSecretKeyID,
config.oauth2AuthURL, config.oauth2TokenURL, config.oauth2RevocationURL, config.oauth2Scopes,
config.apiKeyHeader, config.apiKeyValue, config.apiKeyValueKeyID,
config.customHeaders, config.customHeadersKeyID, userID, now,
)
require.NoError(t, err)
var rowJSON string
err = sqlDB.QueryRowContext(ctx, `SELECT to_jsonb(config) FROM mcp_server_configs AS config WHERE id = $1`, config.id).Scan(&rowJSON)
require.NoError(t, err)
originalJSON[config.id] = rowJSON
}
oauthConfigID := configs[1].id
_, err = sqlDB.ExecContext(ctx, `
INSERT INTO mcp_server_user_tokens (
id, mcp_server_config_id, user_id, access_token, access_token_key_id,
refresh_token, refresh_token_key_id, created_at, updated_at
) VALUES ($1, $2, $3, 'access-token-ciphertext', $4, 'refresh-token-ciphertext', $4, $5, $5)
`, uuid.New(), oauthConfigID, userID, keyDigest, now)
require.NoError(t, err)
chatID := uuid.New()
chatConfigIDs := []uuid.UUID{configs[1].id, configs[0].id}
_, err = sqlDB.ExecContext(ctx, `
INSERT INTO chats (
id, owner_id, organization_id, last_model_config_id, title,
mcp_server_ids, created_at, updated_at
) VALUES ($1, $2, $3, $4, 'Migration 568 Chat', $5, $6, $6)
`, chatID, userID, otherOrgID, modelConfigID, pq.Array(chatConfigIDs), now)
require.NoError(t, err)
version, _, err := next()
require.NoError(t, err)
require.EqualValues(t, 574, version)
var totalConfigs int
err = sqlDB.QueryRowContext(ctx, `SELECT COUNT(*) FROM mcp_server_configs`).Scan(&totalConfigs)
require.NoError(t, err)
require.Equal(t, len(configs), totalConfigs)
for _, config := range configs {
var gotJSON string
var organizationID uuid.UUID
err = sqlDB.QueryRowContext(ctx, `
SELECT to_jsonb(config) - 'organization_id', organization_id
FROM mcp_server_configs AS config
WHERE id = $1
`, config.id).Scan(&gotJSON, &organizationID)
require.NoError(t, err)
require.JSONEq(t, originalJSON[config.id], gotJSON)
require.Equal(t, defaultOrgID, organizationID)
}
var tokenCount int
var tokenConfigID uuid.UUID
err = sqlDB.QueryRowContext(ctx, `SELECT COUNT(*), MIN(mcp_server_config_id::text)::uuid FROM mcp_server_user_tokens`).Scan(&tokenCount, &tokenConfigID)
require.NoError(t, err)
require.Equal(t, 1, tokenCount)
require.Equal(t, oauthConfigID, tokenConfigID)
var slugConstraint string
err = sqlDB.QueryRowContext(ctx, `
SELECT pg_get_constraintdef(oid)
FROM pg_constraint
WHERE conname = 'mcp_server_configs_organization_id_slug_key'
`).Scan(&slugConstraint)
require.NoError(t, err)
require.Equal(t, "UNIQUE (organization_id, slug)", slugConstraint)
getChatIDs := func(t *testing.T, chatID uuid.UUID) []uuid.UUID {
t.Helper()
var ids []uuid.UUID
err := sqlDB.QueryRowContext(ctx, `SELECT mcp_server_ids FROM chats WHERE id = $1`, chatID).Scan(pq.Array(&ids))
require.NoError(t, err)
return ids
}
require.Equal(t, chatConfigIDs, getChatIDs(t, chatID))
// An organization-created config referenced by a chat exercises the down
// sweep: the config is deleted and its chat references removed.
orgLocalConfigID := uuid.New()
_, err = sqlDB.ExecContext(ctx, `
INSERT INTO mcp_server_configs (
id, organization_id, display_name, slug, url, auth_type
) VALUES ($1, $2, 'Org-local config', 'migration-568-org-local', 'https://mcp.example.com/org-local', 'none')
`, orgLocalConfigID, otherOrgID)
require.NoError(t, err)
_, err = sqlDB.ExecContext(ctx, `
UPDATE chats SET mcp_server_ids = $2 WHERE id = $1
`, chatID, pq.Array([]uuid.UUID{orgLocalConfigID}))
require.NoError(t, err)
require.Equal(t, []uuid.UUID{orgLocalConfigID}, getChatIDs(t, chatID))
downSQL, err := os.ReadFile("000574_mcp_server_configs_organization_id.down.sql")
require.NoError(t, err)
_, err = sqlDB.ExecContext(ctx, string(downSQL))
require.NoError(t, err)
err = sqlDB.QueryRowContext(ctx, `SELECT COUNT(*) FROM mcp_server_configs`).Scan(&totalConfigs)
require.NoError(t, err)
require.Equal(t, len(configs), totalConfigs)
require.Empty(t, getChatIDs(t, chatID))
var danglingIDs int
err = sqlDB.QueryRowContext(ctx, `
SELECT COUNT(*)
FROM chats AS chat
CROSS JOIN LATERAL unnest(chat.mcp_server_ids) AS config_id
WHERE NOT EXISTS (SELECT 1 FROM mcp_server_configs WHERE id = config_id)
`).Scan(&danglingIDs)
require.NoError(t, err)
require.Zero(t, danglingIDs)
}
@@ -0,0 +1,149 @@
-- Keeps an MCP config, token, and chat to exercise organization_id and its
-- foreign keys through later migrations and the down sweep. Fixtures run
-- after their matching migration, so this row already has organization_id.
INSERT INTO organizations (
id,
name,
display_name,
description,
icon,
created_at,
updated_at,
is_default,
deleted,
default_org_member_roles
) VALUES (
'f5610000-0000-4000-8000-000000000001',
'fixture-mcp-org',
'Fixture MCP Org',
'',
'',
'2024-01-01 00:00:00+00',
'2024-01-01 00:00:00+00',
FALSE,
FALSE,
'{}'
);
INSERT INTO ai_providers (
id,
type,
name,
display_name,
enabled,
base_url,
created_at,
updated_at
) VALUES (
'f5610000-0000-4000-8000-000000000003',
'openai',
'fixture-mcp-ai-provider',
'Fixture MCP AI Provider',
TRUE,
'https://example.com',
'2024-01-01 00:00:00+00',
'2024-01-01 00:00:00+00'
);
INSERT INTO chat_model_configs (
id,
model,
display_name,
ai_provider_id,
context_limit,
compression_threshold,
created_at,
updated_at
) VALUES (
'f5610000-0000-4000-8000-000000000004',
'fixture-model',
'Fixture Model',
'f5610000-0000-4000-8000-000000000003',
128000,
70,
'2024-01-01 00:00:00+00',
'2024-01-01 00:00:00+00'
);
INSERT INTO dbcrypt_keys (number, active_key_digest, test)
VALUES (561000, 'fixture-000561-key-digest', 'fixture-000561');
-- MCP server config with a ciphertext-shaped secret pair (value + key ID)
-- to assert the backfill leaves secret columns byte-identical. The key ID
-- references dbcrypt_keys(active_key_digest).
INSERT INTO mcp_server_configs (
id,
organization_id,
display_name,
slug,
url,
auth_type,
api_key_value,
api_key_value_key_id,
availability,
enabled,
created_by,
updated_by,
created_at,
updated_at
)
SELECT
'f5610000-0000-4000-8000-000000000005',
(SELECT id FROM organizations WHERE is_default = true LIMIT 1),
'Fixture Org Backfill MCP Server',
'fixture-org-backfill-mcp-server',
'https://mcp.example.com/org-backfill',
'api_key',
'fixture-ciphertext',
'fixture-000561-key-digest',
'default_on',
TRUE,
u.id,
u.id,
'2024-01-01 00:00:00+00',
'2024-01-01 00:00:00+00'
FROM users u
ORDER BY u.created_at, u.id
LIMIT 1;
INSERT INTO mcp_server_user_tokens (
id,
mcp_server_config_id,
user_id,
access_token,
token_type,
created_at,
updated_at
)
SELECT
'f5610000-0000-4000-8000-000000000006',
'f5610000-0000-4000-8000-000000000005',
id,
'fixture-org-backfill-access-token',
'Bearer',
'2024-01-01 00:00:00+00',
'2024-01-01 00:00:00+00'
FROM users
ORDER BY created_at, id
LIMIT 1;
INSERT INTO chats (
id,
owner_id,
organization_id,
last_model_config_id,
title,
mcp_server_ids,
created_at,
updated_at
) VALUES (
'f5610000-0000-4000-8000-000000000007',
(SELECT id FROM users ORDER BY created_at, id LIMIT 1),
'f5610000-0000-4000-8000-000000000001',
'f5610000-0000-4000-8000-000000000004',
'Fixture MCP Org Backfill Chat',
'{f5610000-0000-4000-8000-000000000005}'::uuid[],
'2024-01-01 00:00:00+00',
'2024-01-01 00:00:00+00'
);
+6
View File
@@ -227,6 +227,12 @@ func (c Chat) RBACObject() rbac.Object {
WithGroupACL(c.GroupACL.RBACACL())
}
func (m MCPServerConfig) RBACObject() rbac.Object {
return rbac.ResourceMCPServerConfig.
WithID(m.ID).
InOrg(m.OrganizationID)
}
func (c Chat) IsSubChat() bool {
return c.RootChatID.Valid || c.ParentChatID.Valid
}
+72
View File
@@ -53,6 +53,7 @@ type customQuerier interface {
connectionLogQuerier
aibridgeQuerier
chatQuerier
mcpServerConfigQuerier
}
type templateQuerier interface {
@@ -1199,3 +1200,74 @@ func (q *sqlQuerier) UpdateUserLinkRawJSON(ctx context.Context, userID uuid.UUID
_, err := q.sdb.ExecContext(ctx, "UPDATE user_links SET claims = $2 WHERE user_id = $1", userID, data)
return err
}
type mcpServerConfigQuerier interface {
GetAuthorizedMCPServerConfigs(ctx context.Context, organizationID uuid.UUID, prepared rbac.PreparedAuthorized) ([]MCPServerConfig, error)
}
func (q *sqlQuerier) GetAuthorizedMCPServerConfigs(ctx context.Context, organizationID uuid.UUID, prepared rbac.PreparedAuthorized) ([]MCPServerConfig, error) {
authorizedFilter, err := prepared.CompileToSQL(ctx, regosql.ConvertConfig{
VariableConverter: regosql.MCPServerConfigNoACLConverter(),
})
if err != nil {
return nil, xerrors.Errorf("compile authorized filter: %w", err)
}
filtered, err := insertAuthorizedFilter(getMCPServerConfigsByOrganization, fmt.Sprintf(" AND %s", authorizedFilter))
if err != nil {
return nil, xerrors.Errorf("insert authorized filter: %w", err)
}
// The name comment is for metric tracking
query := fmt.Sprintf("-- name: GetAuthorizedMCPServerConfigs :many\n%s", filtered)
rows, err := q.db.QueryContext(ctx, query, organizationID)
if err != nil {
return nil, err
}
defer rows.Close()
var items []MCPServerConfig
for rows.Next() {
var i MCPServerConfig
if err := rows.Scan(
&i.ID,
&i.DisplayName,
&i.Slug,
&i.Description,
&i.IconURL,
&i.Transport,
&i.Url,
&i.AuthType,
&i.OAuth2ClientID,
&i.OAuth2ClientSecret,
&i.OAuth2ClientSecretKeyID,
&i.OAuth2AuthURL,
&i.OAuth2TokenURL,
&i.OAuth2Scopes,
&i.APIKeyHeader,
&i.APIKeyValue,
&i.APIKeyValueKeyID,
&i.CustomHeaders,
&i.CustomHeadersKeyID,
pq.Array(&i.ToolAllowList),
pq.Array(&i.ToolDenyList),
&i.Availability,
&i.Enabled,
&i.CreatedBy,
&i.UpdatedBy,
&i.CreatedAt,
&i.UpdatedAt,
&i.ModelIntent,
&i.AllowInPlanMode,
&i.ForwardCoderHeaders,
&i.OAuth2RevocationURL,
&i.OrganizationID,
); err != nil {
return nil, err
}
items = append(items, i)
}
if err := rows.Err(); err != nil {
return nil, err
}
return items, nil
}
+17 -1
View File
@@ -465,6 +465,11 @@ const (
ApiKeyScopeWorkspaceBuildOrchestrationDelete APIKeyScope = "workspace_build_orchestration:delete"
ApiKeyScopeWorkspaceBuildOrchestrationRead APIKeyScope = "workspace_build_orchestration:read"
ApiKeyScopeWorkspaceBuildOrchestrationUpdate APIKeyScope = "workspace_build_orchestration:update"
ApiKeyScopeMcpServerConfig APIKeyScope = "mcp_server_config:*"
ApiKeyScopeMcpServerConfigCreate APIKeyScope = "mcp_server_config:create"
ApiKeyScopeMcpServerConfigRead APIKeyScope = "mcp_server_config:read"
ApiKeyScopeMcpServerConfigUpdate APIKeyScope = "mcp_server_config:update"
ApiKeyScopeMcpServerConfigDelete APIKeyScope = "mcp_server_config:delete"
)
func (e *APIKeyScope) Scan(src interface{}) error {
@@ -739,7 +744,12 @@ func (e APIKeyScope) Valid() bool {
ApiKeyScopeWorkspaceBuildOrchestrationCreate,
ApiKeyScopeWorkspaceBuildOrchestrationDelete,
ApiKeyScopeWorkspaceBuildOrchestrationRead,
ApiKeyScopeWorkspaceBuildOrchestrationUpdate:
ApiKeyScopeWorkspaceBuildOrchestrationUpdate,
ApiKeyScopeMcpServerConfig,
ApiKeyScopeMcpServerConfigCreate,
ApiKeyScopeMcpServerConfigRead,
ApiKeyScopeMcpServerConfigUpdate,
ApiKeyScopeMcpServerConfigDelete:
return true
}
return false
@@ -983,6 +993,11 @@ func AllAPIKeyScopeValues() []APIKeyScope {
ApiKeyScopeWorkspaceBuildOrchestrationDelete,
ApiKeyScopeWorkspaceBuildOrchestrationRead,
ApiKeyScopeWorkspaceBuildOrchestrationUpdate,
ApiKeyScopeMcpServerConfig,
ApiKeyScopeMcpServerConfigCreate,
ApiKeyScopeMcpServerConfigRead,
ApiKeyScopeMcpServerConfigUpdate,
ApiKeyScopeMcpServerConfigDelete,
}
}
@@ -5455,6 +5470,7 @@ type MCPServerConfig struct {
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"`
OrganizationID uuid.UUID `db:"organization_id" json:"organization_id"`
}
type MCPServerUserToken struct {
+7 -5
View File
@@ -161,6 +161,7 @@ type sqlcQuerier interface {
DeleteLicense(ctx context.Context, id int32) (int32, error)
DeleteMCPServerConfigByID(ctx context.Context, id uuid.UUID) error
DeleteMCPServerUserToken(ctx context.Context, arg DeleteMCPServerUserTokenParams) error
DeleteMCPServerUserTokensByConfigID(ctx context.Context, mcpServerConfigID uuid.UUID) error
DeleteOAuth2ProviderAppByClientID(ctx context.Context, id uuid.UUID) error
DeleteOAuth2ProviderAppByID(ctx context.Context, id uuid.UUID) error
DeleteOAuth2ProviderAppCodeByID(ctx context.Context, id uuid.UUID) error
@@ -554,7 +555,8 @@ type sqlcQuerier interface {
// Check both to ensure the selected config is actually usable.
GetEnabledChatModelConfigByID(ctx context.Context, id uuid.UUID) (ChatModelConfig, error)
GetEnabledChatModelConfigs(ctx context.Context) ([]GetEnabledChatModelConfigsRow, error)
GetEnabledMCPServerConfigs(ctx context.Context) ([]MCPServerConfig, error)
GetEnabledMCPServerConfigsByOrganization(ctx context.Context, organizationID uuid.UUID) ([]MCPServerConfig, error)
GetEnabledMCPServerConfigsByOrganizationAndIDs(ctx context.Context, arg GetEnabledMCPServerConfigsByOrganizationAndIDsParams) ([]MCPServerConfig, error)
// GetExternalAgentTokensByTemplateID returns the auth tokens for all
// non-deleted external agents on the latest build of every running workspace
// of the given template. "Running" means the latest build has
@@ -579,7 +581,7 @@ type sqlcQuerier interface {
// param created_at_opt: The created_at timestamp to filter by. This parameter is usd for pagination - it fetches notifications created before the specified timestamp if it is not the zero value
// param limit_opt: The limit of notifications to fetch. If the limit is not specified, it defaults to 25
GetFilteredInboxNotificationsByUserID(ctx context.Context, arg GetFilteredInboxNotificationsByUserIDParams) ([]InboxNotification, error)
GetForcedMCPServerConfigs(ctx context.Context) ([]MCPServerConfig, error)
GetForcedMCPServerConfigsByOrganization(ctx context.Context, organizationID uuid.UUID) ([]MCPServerConfig, error)
GetGitSSHKey(ctx context.Context, userID uuid.UUID) (GitSSHKey, error)
GetGroupAIBudget(ctx context.Context, groupID uuid.UUID) (GroupAIBudget, error)
GetGroupByID(ctx context.Context, id uuid.UUID) (Group, error)
@@ -645,9 +647,9 @@ type sqlcQuerier interface {
GetLicenses(ctx context.Context) ([]License, error)
GetLogoURL(ctx context.Context) (string, error)
GetMCPServerConfigByID(ctx context.Context, id uuid.UUID) (MCPServerConfig, error)
GetMCPServerConfigBySlug(ctx context.Context, slug string) (MCPServerConfig, error)
GetMCPServerConfigs(ctx context.Context) ([]MCPServerConfig, error)
GetMCPServerConfigsByIDs(ctx context.Context, ids []uuid.UUID) ([]MCPServerConfig, error)
GetMCPServerConfigByIDForUpdate(ctx context.Context, id uuid.UUID) (MCPServerConfig, error)
GetMCPServerConfigByOrganizationAndSlug(ctx context.Context, arg GetMCPServerConfigByOrganizationAndSlugParams) (MCPServerConfig, error)
GetMCPServerConfigsByOrganization(ctx context.Context, organizationID uuid.UUID) ([]MCPServerConfig, error)
GetMCPServerUserToken(ctx context.Context, arg GetMCPServerUserTokenParams) (MCPServerUserToken, error)
GetMCPServerUserTokensByUserID(ctx context.Context, userID uuid.UUID) ([]MCPServerUserToken, error)
// Must be called from within a transaction. The row lock is released
+194 -97
View File
@@ -17133,19 +17133,32 @@ func (q *sqlQuerier) DeleteMCPServerUserToken(ctx context.Context, arg DeleteMCP
return err
}
const getEnabledMCPServerConfigs = `-- name: GetEnabledMCPServerConfigs :many
const deleteMCPServerUserTokensByConfigID = `-- name: DeleteMCPServerUserTokensByConfigID :exec
DELETE FROM
mcp_server_user_tokens
WHERE
mcp_server_config_id = $1::uuid
`
func (q *sqlQuerier) DeleteMCPServerUserTokensByConfigID(ctx context.Context, mcpServerConfigID uuid.UUID) error {
_, err := q.db.ExecContext(ctx, deleteMCPServerUserTokensByConfigID, mcpServerConfigID)
return err
}
const getEnabledMCPServerConfigsByOrganization = `-- name: GetEnabledMCPServerConfigsByOrganization :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, oauth2_revocation_url
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, organization_id
FROM
mcp_server_configs
WHERE
enabled = TRUE
organization_id = $1::uuid
AND enabled = TRUE
ORDER BY
display_name ASC
`
func (q *sqlQuerier) GetEnabledMCPServerConfigs(ctx context.Context) ([]MCPServerConfig, error) {
rows, err := q.db.QueryContext(ctx, getEnabledMCPServerConfigs)
func (q *sqlQuerier) GetEnabledMCPServerConfigsByOrganization(ctx context.Context, organizationID uuid.UUID) ([]MCPServerConfig, error) {
rows, err := q.db.QueryContext(ctx, getEnabledMCPServerConfigsByOrganization, organizationID)
if err != nil {
return nil, err
}
@@ -17185,6 +17198,7 @@ func (q *sqlQuerier) GetEnabledMCPServerConfigs(ctx context.Context) ([]MCPServe
&i.AllowInPlanMode,
&i.ForwardCoderHeaders,
&i.OAuth2RevocationURL,
&i.OrganizationID,
); err != nil {
return nil, err
}
@@ -17199,20 +17213,26 @@ func (q *sqlQuerier) GetEnabledMCPServerConfigs(ctx context.Context) ([]MCPServe
return items, nil
}
const getForcedMCPServerConfigs = `-- name: GetForcedMCPServerConfigs :many
const getEnabledMCPServerConfigsByOrganizationAndIDs = `-- name: GetEnabledMCPServerConfigsByOrganizationAndIDs :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, oauth2_revocation_url
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, organization_id
FROM
mcp_server_configs
WHERE
enabled = TRUE
AND availability = 'force_on'
organization_id = $1::uuid
AND id = ANY($2::uuid[])
AND enabled = TRUE
ORDER BY
display_name ASC
`
func (q *sqlQuerier) GetForcedMCPServerConfigs(ctx context.Context) ([]MCPServerConfig, error) {
rows, err := q.db.QueryContext(ctx, getForcedMCPServerConfigs)
type GetEnabledMCPServerConfigsByOrganizationAndIDsParams struct {
OrganizationID uuid.UUID `db:"organization_id" json:"organization_id"`
IDs []uuid.UUID `db:"ids" json:"ids"`
}
func (q *sqlQuerier) GetEnabledMCPServerConfigsByOrganizationAndIDs(ctx context.Context, arg GetEnabledMCPServerConfigsByOrganizationAndIDsParams) ([]MCPServerConfig, error) {
rows, err := q.db.QueryContext(ctx, getEnabledMCPServerConfigsByOrganizationAndIDs, arg.OrganizationID, pq.Array(arg.IDs))
if err != nil {
return nil, err
}
@@ -17252,6 +17272,76 @@ func (q *sqlQuerier) GetForcedMCPServerConfigs(ctx context.Context) ([]MCPServer
&i.AllowInPlanMode,
&i.ForwardCoderHeaders,
&i.OAuth2RevocationURL,
&i.OrganizationID,
); err != nil {
return nil, err
}
items = append(items, i)
}
if err := rows.Close(); err != nil {
return nil, err
}
if err := rows.Err(); err != nil {
return nil, err
}
return items, nil
}
const getForcedMCPServerConfigsByOrganization = `-- name: GetForcedMCPServerConfigsByOrganization :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, oauth2_revocation_url, organization_id
FROM
mcp_server_configs
WHERE
organization_id = $1::uuid
AND enabled = TRUE
AND availability = 'force_on'
ORDER BY
display_name ASC
`
func (q *sqlQuerier) GetForcedMCPServerConfigsByOrganization(ctx context.Context, organizationID uuid.UUID) ([]MCPServerConfig, error) {
rows, err := q.db.QueryContext(ctx, getForcedMCPServerConfigsByOrganization, organizationID)
if err != nil {
return nil, err
}
defer rows.Close()
var items []MCPServerConfig
for rows.Next() {
var i MCPServerConfig
if err := rows.Scan(
&i.ID,
&i.DisplayName,
&i.Slug,
&i.Description,
&i.IconURL,
&i.Transport,
&i.Url,
&i.AuthType,
&i.OAuth2ClientID,
&i.OAuth2ClientSecret,
&i.OAuth2ClientSecretKeyID,
&i.OAuth2AuthURL,
&i.OAuth2TokenURL,
&i.OAuth2Scopes,
&i.APIKeyHeader,
&i.APIKeyValue,
&i.APIKeyValueKeyID,
&i.CustomHeaders,
&i.CustomHeadersKeyID,
pq.Array(&i.ToolAllowList),
pq.Array(&i.ToolDenyList),
&i.Availability,
&i.Enabled,
&i.CreatedBy,
&i.UpdatedBy,
&i.CreatedAt,
&i.UpdatedAt,
&i.ModelIntent,
&i.AllowInPlanMode,
&i.ForwardCoderHeaders,
&i.OAuth2RevocationURL,
&i.OrganizationID,
); err != nil {
return nil, err
}
@@ -17268,7 +17358,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, oauth2_revocation_url
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, organization_id
FROM
mcp_server_configs
WHERE
@@ -17310,21 +17400,23 @@ func (q *sqlQuerier) GetMCPServerConfigByID(ctx context.Context, id uuid.UUID) (
&i.AllowInPlanMode,
&i.ForwardCoderHeaders,
&i.OAuth2RevocationURL,
&i.OrganizationID,
)
return i, err
}
const getMCPServerConfigBySlug = `-- name: GetMCPServerConfigBySlug :one
const getMCPServerConfigByIDForUpdate = `-- name: GetMCPServerConfigByIDForUpdate :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, oauth2_revocation_url
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, organization_id
FROM
mcp_server_configs
WHERE
slug = $1::text
id = $1::uuid
FOR UPDATE
`
func (q *sqlQuerier) GetMCPServerConfigBySlug(ctx context.Context, slug string) (MCPServerConfig, error) {
row := q.db.QueryRowContext(ctx, getMCPServerConfigBySlug, slug)
func (q *sqlQuerier) GetMCPServerConfigByIDForUpdate(ctx context.Context, id uuid.UUID) (MCPServerConfig, error) {
row := q.db.QueryRowContext(ctx, getMCPServerConfigByIDForUpdate, id)
var i MCPServerConfig
err := row.Scan(
&i.ID,
@@ -17358,87 +17450,81 @@ func (q *sqlQuerier) GetMCPServerConfigBySlug(ctx context.Context, slug string)
&i.AllowInPlanMode,
&i.ForwardCoderHeaders,
&i.OAuth2RevocationURL,
&i.OrganizationID,
)
return i, err
}
const getMCPServerConfigs = `-- name: GetMCPServerConfigs :many
const getMCPServerConfigByOrganizationAndSlug = `-- name: GetMCPServerConfigByOrganizationAndSlug :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, oauth2_revocation_url
FROM
mcp_server_configs
ORDER BY
display_name ASC
`
func (q *sqlQuerier) GetMCPServerConfigs(ctx context.Context) ([]MCPServerConfig, error) {
rows, err := q.db.QueryContext(ctx, getMCPServerConfigs)
if err != nil {
return nil, err
}
defer rows.Close()
var items []MCPServerConfig
for rows.Next() {
var i MCPServerConfig
if err := rows.Scan(
&i.ID,
&i.DisplayName,
&i.Slug,
&i.Description,
&i.IconURL,
&i.Transport,
&i.Url,
&i.AuthType,
&i.OAuth2ClientID,
&i.OAuth2ClientSecret,
&i.OAuth2ClientSecretKeyID,
&i.OAuth2AuthURL,
&i.OAuth2TokenURL,
&i.OAuth2Scopes,
&i.APIKeyHeader,
&i.APIKeyValue,
&i.APIKeyValueKeyID,
&i.CustomHeaders,
&i.CustomHeadersKeyID,
pq.Array(&i.ToolAllowList),
pq.Array(&i.ToolDenyList),
&i.Availability,
&i.Enabled,
&i.CreatedBy,
&i.UpdatedBy,
&i.CreatedAt,
&i.UpdatedAt,
&i.ModelIntent,
&i.AllowInPlanMode,
&i.ForwardCoderHeaders,
&i.OAuth2RevocationURL,
); err != nil {
return nil, err
}
items = append(items, i)
}
if err := rows.Close(); err != nil {
return nil, err
}
if err := rows.Err(); err != nil {
return nil, err
}
return items, nil
}
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, oauth2_revocation_url
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, organization_id
FROM
mcp_server_configs
WHERE
id = ANY($1::uuid[])
organization_id = $1::uuid
AND slug = $2::text
`
type GetMCPServerConfigByOrganizationAndSlugParams struct {
OrganizationID uuid.UUID `db:"organization_id" json:"organization_id"`
Slug string `db:"slug" json:"slug"`
}
func (q *sqlQuerier) GetMCPServerConfigByOrganizationAndSlug(ctx context.Context, arg GetMCPServerConfigByOrganizationAndSlugParams) (MCPServerConfig, error) {
row := q.db.QueryRowContext(ctx, getMCPServerConfigByOrganizationAndSlug, arg.OrganizationID, arg.Slug)
var i MCPServerConfig
err := row.Scan(
&i.ID,
&i.DisplayName,
&i.Slug,
&i.Description,
&i.IconURL,
&i.Transport,
&i.Url,
&i.AuthType,
&i.OAuth2ClientID,
&i.OAuth2ClientSecret,
&i.OAuth2ClientSecretKeyID,
&i.OAuth2AuthURL,
&i.OAuth2TokenURL,
&i.OAuth2Scopes,
&i.APIKeyHeader,
&i.APIKeyValue,
&i.APIKeyValueKeyID,
&i.CustomHeaders,
&i.CustomHeadersKeyID,
pq.Array(&i.ToolAllowList),
pq.Array(&i.ToolDenyList),
&i.Availability,
&i.Enabled,
&i.CreatedBy,
&i.UpdatedBy,
&i.CreatedAt,
&i.UpdatedAt,
&i.ModelIntent,
&i.AllowInPlanMode,
&i.ForwardCoderHeaders,
&i.OAuth2RevocationURL,
&i.OrganizationID,
)
return i, err
}
const getMCPServerConfigsByOrganization = `-- name: GetMCPServerConfigsByOrganization :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, oauth2_revocation_url, organization_id
FROM
mcp_server_configs
WHERE
organization_id = $1::uuid
-- Authorize Filter clause will be injected below in GetAuthorizedMCPServerConfigs
-- @authorize_filter
ORDER BY
display_name ASC
`
func (q *sqlQuerier) GetMCPServerConfigsByIDs(ctx context.Context, ids []uuid.UUID) ([]MCPServerConfig, error) {
rows, err := q.db.QueryContext(ctx, getMCPServerConfigsByIDs, pq.Array(ids))
func (q *sqlQuerier) GetMCPServerConfigsByOrganization(ctx context.Context, organizationID uuid.UUID) ([]MCPServerConfig, error) {
rows, err := q.db.QueryContext(ctx, getMCPServerConfigsByOrganization, organizationID)
if err != nil {
return nil, err
}
@@ -17478,6 +17564,7 @@ func (q *sqlQuerier) GetMCPServerConfigsByIDs(ctx context.Context, ids []uuid.UU
&i.AllowInPlanMode,
&i.ForwardCoderHeaders,
&i.OAuth2RevocationURL,
&i.OrganizationID,
); err != nil {
return nil, err
}
@@ -17574,6 +17661,8 @@ func (q *sqlQuerier) GetMCPServerUserTokensByUserID(ctx context.Context, userID
const insertMCPServerConfig = `-- name: InsertMCPServerConfig :one
INSERT INTO mcp_server_configs (
id,
organization_id,
display_name,
slug,
description,
@@ -17603,8 +17692,8 @@ INSERT INTO mcp_server_configs (
created_by,
updated_by
) VALUES (
$1::text,
$2::text,
$1::uuid,
$2::uuid,
$3::text,
$4::text,
$5::text,
@@ -17622,21 +17711,25 @@ INSERT INTO mcp_server_configs (
$17::text,
$18::text,
$19::text,
$20::text[],
$21::text[],
$22::text,
$23::boolean,
$24::boolean,
$20::text,
$21::text,
$22::text[],
$23::text[],
$24::text,
$25::boolean,
$26::boolean,
$27::uuid,
$28::uuid
$27::boolean,
$28::boolean,
$29::uuid,
$30::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, oauth2_revocation_url
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, organization_id
`
type InsertMCPServerConfigParams struct {
ID uuid.UUID `db:"id" json:"id"`
OrganizationID uuid.UUID `db:"organization_id" json:"organization_id"`
DisplayName string `db:"display_name" json:"display_name"`
Slug string `db:"slug" json:"slug"`
Description string `db:"description" json:"description"`
@@ -17669,6 +17762,8 @@ type InsertMCPServerConfigParams struct {
func (q *sqlQuerier) InsertMCPServerConfig(ctx context.Context, arg InsertMCPServerConfigParams) (MCPServerConfig, error) {
row := q.db.QueryRowContext(ctx, insertMCPServerConfig,
arg.ID,
arg.OrganizationID,
arg.DisplayName,
arg.Slug,
arg.Description,
@@ -17731,6 +17826,7 @@ func (q *sqlQuerier) InsertMCPServerConfig(ctx context.Context, arg InsertMCPSer
&i.AllowInPlanMode,
&i.ForwardCoderHeaders,
&i.OAuth2RevocationURL,
&i.OrganizationID,
)
return i, err
}
@@ -17818,7 +17914,7 @@ SET
WHERE
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, oauth2_revocation_url
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, organization_id
`
type UpdateMCPServerConfigParams struct {
@@ -17916,6 +18012,7 @@ func (q *sqlQuerier) UpdateMCPServerConfig(ctx context.Context, arg UpdateMCPSer
&i.AllowInPlanMode,
&i.ForwardCoderHeaders,
&i.OAuth2RevocationURL,
&i.OrganizationID,
)
return i, err
}
+37 -9
View File
@@ -6,55 +6,75 @@ FROM
WHERE
id = @id::uuid;
-- name: GetMCPServerConfigBySlug :one
-- name: GetMCPServerConfigByIDForUpdate :one
SELECT
*
FROM
mcp_server_configs
WHERE
slug = @slug::text;
id = @id::uuid
FOR UPDATE;
-- name: GetMCPServerConfigs :many
-- name: GetMCPServerConfigByOrganizationAndSlug :one
SELECT
*
FROM
mcp_server_configs
WHERE
organization_id = @organization_id::uuid
AND slug = @slug::text;
-- name: GetMCPServerConfigsByOrganization :many
SELECT
*
FROM
mcp_server_configs
WHERE
organization_id = @organization_id::uuid
-- Authorize Filter clause will be injected below in GetAuthorizedMCPServerConfigs
-- @authorize_filter
ORDER BY
display_name ASC;
-- name: GetEnabledMCPServerConfigs :many
-- name: GetEnabledMCPServerConfigsByOrganization :many
SELECT
*
FROM
mcp_server_configs
WHERE
enabled = TRUE
organization_id = @organization_id::uuid
AND enabled = TRUE
ORDER BY
display_name ASC;
-- name: GetMCPServerConfigsByIDs :many
-- name: GetEnabledMCPServerConfigsByOrganizationAndIDs :many
SELECT
*
FROM
mcp_server_configs
WHERE
id = ANY(@ids::uuid[])
organization_id = @organization_id::uuid
AND id = ANY(@ids::uuid[])
AND enabled = TRUE
ORDER BY
display_name ASC;
-- name: GetForcedMCPServerConfigs :many
-- name: GetForcedMCPServerConfigsByOrganization :many
SELECT
*
FROM
mcp_server_configs
WHERE
enabled = TRUE
organization_id = @organization_id::uuid
AND enabled = TRUE
AND availability = 'force_on'
ORDER BY
display_name ASC;
-- name: InsertMCPServerConfig :one
INSERT INTO mcp_server_configs (
id,
organization_id,
display_name,
slug,
description,
@@ -84,6 +104,8 @@ INSERT INTO mcp_server_configs (
created_by,
updated_by
) VALUES (
@id::uuid,
@organization_id::uuid,
@display_name::text,
@slug::text,
@description::text,
@@ -257,6 +279,12 @@ WHERE
mcp_server_config_id = @mcp_server_config_id::uuid
AND user_id = @user_id::uuid;
-- name: DeleteMCPServerUserTokensByConfigID :exec
DELETE FROM
mcp_server_user_tokens
WHERE
mcp_server_config_id = @mcp_server_config_id::uuid;
-- name: CleanupDeletedMCPServerIDsFromChats :exec
UPDATE chats
SET mcp_server_ids = (
+1 -1
View File
@@ -53,8 +53,8 @@ const (
UniqueJfrogXrayScansPkey UniqueConstraint = "jfrog_xray_scans_pkey" // ALTER TABLE ONLY jfrog_xray_scans ADD CONSTRAINT jfrog_xray_scans_pkey PRIMARY KEY (agent_id, workspace_id);
UniqueLicensesJWTKey UniqueConstraint = "licenses_jwt_key" // ALTER TABLE ONLY licenses ADD CONSTRAINT licenses_jwt_key UNIQUE (jwt);
UniqueLicensesPkey UniqueConstraint = "licenses_pkey" // ALTER TABLE ONLY licenses ADD CONSTRAINT licenses_pkey PRIMARY KEY (id);
UniqueMcpServerConfigsOrganizationIDSlugKey UniqueConstraint = "mcp_server_configs_organization_id_slug_key" // ALTER TABLE ONLY mcp_server_configs ADD CONSTRAINT mcp_server_configs_organization_id_slug_key UNIQUE (organization_id, slug);
UniqueMcpServerConfigsPkey UniqueConstraint = "mcp_server_configs_pkey" // ALTER TABLE ONLY mcp_server_configs ADD CONSTRAINT mcp_server_configs_pkey PRIMARY KEY (id);
UniqueMcpServerConfigsSlugKey UniqueConstraint = "mcp_server_configs_slug_key" // ALTER TABLE ONLY mcp_server_configs ADD CONSTRAINT mcp_server_configs_slug_key UNIQUE (slug);
UniqueMcpServerUserTokensMcpServerConfigIDUserIDKey UniqueConstraint = "mcp_server_user_tokens_mcp_server_config_id_user_id_key" // ALTER TABLE ONLY mcp_server_user_tokens ADD CONSTRAINT mcp_server_user_tokens_mcp_server_config_id_user_id_key UNIQUE (mcp_server_config_id, user_id);
UniqueMcpServerUserTokensPkey UniqueConstraint = "mcp_server_user_tokens_pkey" // ALTER TABLE ONLY mcp_server_user_tokens ADD CONSTRAINT mcp_server_user_tokens_pkey PRIMARY KEY (id);
UniqueNotificationMessagesPkey UniqueConstraint = "notification_messages_pkey" // ALTER TABLE ONLY notification_messages ADD CONSTRAINT notification_messages_pkey PRIMARY KEY (id);
+82 -46
View File
@@ -1244,6 +1244,57 @@ func (api *API) validateExplicitChatModelConfigAvailable(
return status, resp
}
func validateChatMCPServerIDs(
ctx context.Context,
db database.Store,
organizationID uuid.UUID,
ids []uuid.UUID,
) (normalized []uuid.UUID, invalid []uuid.UUID, err error) {
unique := make([]uuid.UUID, 0, len(ids))
seen := make(map[uuid.UUID]struct{}, len(ids))
for _, id := range ids {
if _, ok := seen[id]; ok {
continue
}
seen[id] = struct{}{}
unique = append(unique, id)
}
if len(unique) == 0 {
return unique, nil, nil
}
configs, err := db.GetEnabledMCPServerConfigsByOrganizationAndIDs(ctx, database.GetEnabledMCPServerConfigsByOrganizationAndIDsParams{
OrganizationID: organizationID,
IDs: unique,
})
if err != nil {
return nil, nil, xerrors.Errorf("get enabled MCP server configs for organization: %w", err)
}
valid := make(map[uuid.UUID]struct{}, len(configs))
for _, config := range configs {
valid[config.ID] = struct{}{}
}
invalid = make([]uuid.UUID, 0, len(unique)-len(valid))
for _, id := range unique {
if _, ok := valid[id]; !ok {
invalid = append(invalid, id)
}
}
return unique, invalid, nil
}
func invalidChatMCPServerIDsResponse(ids []uuid.UUID) codersdk.Response {
invalid := make([]string, 0, len(ids))
for _, id := range ids {
invalid = append(invalid, id.String())
}
return codersdk.Response{
Message: "One or more MCP server IDs are invalid or disabled.",
Detail: fmt.Sprintf("Invalid IDs: %s", strings.Join(invalid, ", ")),
}
}
// EXPERIMENTAL: this endpoint is experimental and is subject to change.
//
// @Summary Create chat
@@ -1337,34 +1388,18 @@ func (api *API) postChats(rw http.ResponseWriter, r *http.Request) {
return
}
// Validate MCP server IDs exist.
if len(req.MCPServerIDs) > 0 {
//nolint:gocritic // Need to validate MCP server IDs exist.
existingConfigs, err := api.Database.GetMCPServerConfigsByIDs(dbauthz.AsSystemRestricted(ctx), req.MCPServerIDs)
if err != nil {
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
Message: "Failed to validate MCP server IDs.",
Detail: err.Error(),
})
return
}
if len(existingConfigs) != len(req.MCPServerIDs) {
found := make(map[uuid.UUID]struct{}, len(existingConfigs))
for _, c := range existingConfigs {
found[c.ID] = struct{}{}
}
var missing []string
for _, id := range req.MCPServerIDs {
if _, ok := found[id]; !ok {
missing = append(missing, id.String())
}
}
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
Message: "One or more MCP server IDs are invalid.",
Detail: fmt.Sprintf("Invalid IDs: %s", strings.Join(missing, ", ")),
})
return
}
normalizedMCPServerIDs, invalidMCPServerIDs, err := validateChatMCPServerIDs(ctx, api.Database, req.OrganizationID, req.MCPServerIDs)
if err != nil {
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
Message: "Failed to validate MCP server IDs.",
Detail: err.Error(),
})
return
}
req.MCPServerIDs = normalizedMCPServerIDs
if len(invalidMCPServerIDs) > 0 {
httpapi.Write(ctx, rw, http.StatusBadRequest, invalidChatMCPServerIDsResponse(invalidMCPServerIDs))
return
}
mcpServerIDs := req.MCPServerIDs
@@ -2735,10 +2770,8 @@ func (api *API) postChatMessages(rw http.ResponseWriter, r *http.Request) {
return
}
// Validate MCP server IDs exist.
if req.MCPServerIDs != nil && len(*req.MCPServerIDs) > 0 {
//nolint:gocritic // Need to validate MCP server IDs exist.
existingConfigs, err := api.Database.GetMCPServerConfigsByIDs(dbauthz.AsSystemRestricted(ctx), *req.MCPServerIDs)
if req.MCPServerIDs != nil {
normalizedMCPServerIDs, invalidMCPServerIDs, err := validateChatMCPServerIDs(ctx, api.Database, chat.OrganizationID, *req.MCPServerIDs)
if err != nil {
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
Message: "Failed to validate MCP server IDs.",
@@ -2746,21 +2779,24 @@ func (api *API) postChatMessages(rw http.ResponseWriter, r *http.Request) {
})
return
}
if len(existingConfigs) != len(*req.MCPServerIDs) {
found := make(map[uuid.UUID]struct{}, len(existingConfigs))
for _, c := range existingConfigs {
found[c.ID] = struct{}{}
req.MCPServerIDs = &normalizedMCPServerIDs
// IDs already persisted on the chat are exempt: a server that
// is disabled or revoked after selection must not block sends.
// The generation path skips servers the chat can no longer use,
// and keeping the ID preserves the selection if the server is
// re-enabled.
persisted := make(map[uuid.UUID]struct{}, len(chat.MCPServerIDs))
for _, id := range chat.MCPServerIDs {
persisted[id] = struct{}{}
}
newlyInvalid := make([]uuid.UUID, 0, len(invalidMCPServerIDs))
for _, id := range invalidMCPServerIDs {
if _, ok := persisted[id]; !ok {
newlyInvalid = append(newlyInvalid, id)
}
var missing []string
for _, id := range *req.MCPServerIDs {
if _, ok := found[id]; !ok {
missing = append(missing, id.String())
}
}
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
Message: "One or more MCP server IDs are invalid.",
Detail: fmt.Sprintf("Invalid IDs: %s", strings.Join(missing, ", ")),
})
}
if len(newlyInvalid) > 0 {
httpapi.Write(ctx, rw, http.StatusBadRequest, invalidChatMCPServerIDsResponse(newlyInvalid))
return
}
}
+244 -1
View File
@@ -570,6 +570,249 @@ func TestPostChats(t *testing.T) {
}))
})
t.Run("MCPServerIDsCrossOrgRejected", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
client, db := newChatClientWithDatabase(t)
firstUser := coderdtest.CreateFirstUser(t, client.Client)
_ = createChatModelConfig(t, client)
defaultOrgConfig := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{
OrganizationID: firstUser.OrganizationID,
Enabled: true,
})
secondOrg := dbgen.Organization(t, db, database.Organization{})
memberClientRaw, _ := coderdtest.CreateAnotherUser(t, client.Client, secondOrg.ID, rbac.ScopedRoleAgentsAccess(secondOrg.ID))
memberClient := codersdk.NewExperimentalClient(memberClientRaw)
_, err := memberClient.CreateChat(ctx, codersdk.CreateChatRequest{
OrganizationID: secondOrg.ID,
Content: []codersdk.ChatInputPart{
{
Type: codersdk.ChatInputPartTypeText,
Text: "chat with a default-org MCP server",
},
},
MCPServerIDs: []uuid.UUID{defaultOrgConfig.ID},
})
sdkErr := requireSDKError(t, err, http.StatusBadRequest)
require.Equal(t, "One or more MCP server IDs are invalid or disabled.", sdkErr.Message)
require.Equal(t, "Invalid IDs: "+defaultOrgConfig.ID.String(), sdkErr.Detail)
})
t.Run("MCPServerIDsDuplicatesNormalized", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
client, db := newChatClientWithDatabase(t)
coderdtest.CreateFirstUser(t, client.Client)
_ = createChatModelConfig(t, client)
secondOrg := dbgen.Organization(t, db, database.Organization{})
orgConfig := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{
OrganizationID: secondOrg.ID,
Enabled: true,
})
memberClientRaw, _ := coderdtest.CreateAnotherUser(t, client.Client, secondOrg.ID, rbac.ScopedRoleAgentsAccess(secondOrg.ID))
memberClient := codersdk.NewExperimentalClient(memberClientRaw)
chat, err := memberClient.CreateChat(ctx, codersdk.CreateChatRequest{
OrganizationID: secondOrg.ID,
Content: []codersdk.ChatInputPart{
{
Type: codersdk.ChatInputPartTypeText,
Text: "chat with a duplicated MCP server ID",
},
},
MCPServerIDs: []uuid.UUID{orgConfig.ID, orgConfig.ID},
})
require.NoError(t, err)
require.Equal(t, []uuid.UUID{orgConfig.ID}, chat.MCPServerIDs)
_, err = memberClient.CreateChatMessage(ctx, chat.ID, codersdk.CreateChatMessageRequest{
Content: []codersdk.ChatInputPart{
{
Type: codersdk.ChatInputPartTypeText,
Text: "update to a duplicated MCP server ID",
},
},
MCPServerIDs: &[]uuid.UUID{orgConfig.ID, orgConfig.ID},
})
require.NoError(t, err)
storedChat, err := db.GetChatByID(dbauthz.AsSystemRestricted(ctx), chat.ID)
require.NoError(t, err)
require.Equal(t, []uuid.UUID{orgConfig.ID}, storedChat.MCPServerIDs)
})
t.Run("MCPServerIDsDisabledConfigRejected", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
client, db := newChatClientWithDatabase(t)
firstUser := coderdtest.CreateFirstUser(t, client.Client)
_ = createChatModelConfig(t, client)
enabledCfg := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{
OrganizationID: firstUser.OrganizationID,
Enabled: true,
})
disabledCfg, err := client.Client.UpdateMCPServerConfig(ctx, enabledCfg.OrganizationID, enabledCfg.ID, codersdk.UpdateMCPServerConfigRequest{
Enabled: ptr.Ref(false),
})
require.NoError(t, err)
memberClientRaw, _ := coderdtest.CreateAnotherUser(t, client.Client, firstUser.OrganizationID, rbac.ScopedRoleAgentsAccess(firstUser.OrganizationID))
memberClient := codersdk.NewExperimentalClient(memberClientRaw)
_, err = memberClient.CreateChat(ctx, codersdk.CreateChatRequest{
OrganizationID: firstUser.OrganizationID,
Content: []codersdk.ChatInputPart{
{
Type: codersdk.ChatInputPartTypeText,
Text: "chat with a disabled MCP server",
},
},
MCPServerIDs: []uuid.UUID{disabledCfg.ID},
})
sdkErr := requireSDKError(t, err, http.StatusBadRequest)
require.Equal(t, "One or more MCP server IDs are invalid or disabled.", sdkErr.Message)
require.Equal(t, "Invalid IDs: "+disabledCfg.ID.String(), sdkErr.Detail)
chat, err := memberClient.CreateChat(ctx, codersdk.CreateChatRequest{
OrganizationID: firstUser.OrganizationID,
Content: []codersdk.ChatInputPart{
{
Type: codersdk.ChatInputPartTypeText,
Text: "chat without MCP servers",
},
},
})
require.NoError(t, err)
_, err = memberClient.CreateChatMessage(ctx, chat.ID, codersdk.CreateChatMessageRequest{
Content: []codersdk.ChatInputPart{
{
Type: codersdk.ChatInputPartTypeText,
Text: "selecting the disabled config",
},
},
MCPServerIDs: &[]uuid.UUID{disabledCfg.ID},
})
sdkErr = requireSDKError(t, err, http.StatusBadRequest)
require.Equal(t, "One or more MCP server IDs are invalid or disabled.", sdkErr.Message)
require.Equal(t, "Invalid IDs: "+disabledCfg.ID.String(), sdkErr.Detail)
})
t.Run("MCPServerIDsPersistedDisabledStillSendable", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
client, db := newChatClientWithDatabase(t)
firstUser := coderdtest.CreateFirstUser(t, client.Client)
_ = createChatModelConfig(t, client)
cfg := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{
OrganizationID: firstUser.OrganizationID,
Enabled: true,
})
memberClientRaw, _ := coderdtest.CreateAnotherUser(t, client.Client, firstUser.OrganizationID, rbac.ScopedRoleAgentsAccess(firstUser.OrganizationID))
memberClient := codersdk.NewExperimentalClient(memberClientRaw)
chat, err := memberClient.CreateChat(ctx, codersdk.CreateChatRequest{
OrganizationID: firstUser.OrganizationID,
Content: []codersdk.ChatInputPart{
{
Type: codersdk.ChatInputPartTypeText,
Text: "chat with a server that is disabled later",
},
},
MCPServerIDs: []uuid.UUID{cfg.ID},
})
require.NoError(t, err)
_, err = client.Client.UpdateMCPServerConfig(ctx, cfg.OrganizationID, cfg.ID, codersdk.UpdateMCPServerConfigRequest{
Enabled: ptr.Ref(false),
})
require.NoError(t, err)
// The frontend resubmits the persisted selection on every
// send, so a server disabled after selection must not block
// the send.
_, err = memberClient.CreateChatMessage(ctx, chat.ID, codersdk.CreateChatMessageRequest{
Content: []codersdk.ChatInputPart{
{
Type: codersdk.ChatInputPartTypeText,
Text: "still sendable after the server was disabled",
},
},
MCPServerIDs: &[]uuid.UUID{cfg.ID},
})
require.NoError(t, err)
storedChat, err := db.GetChatByID(dbauthz.AsSystemRestricted(ctx), chat.ID)
require.NoError(t, err)
require.Equal(t, []uuid.UUID{cfg.ID}, storedChat.MCPServerIDs)
secondCfg := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{
OrganizationID: firstUser.OrganizationID,
Enabled: true,
})
_, err = client.Client.UpdateMCPServerConfig(ctx, secondCfg.OrganizationID, secondCfg.ID, codersdk.UpdateMCPServerConfigRequest{
Enabled: ptr.Ref(false),
})
require.NoError(t, err)
_, err = memberClient.CreateChatMessage(ctx, chat.ID, codersdk.CreateChatMessageRequest{
Content: []codersdk.ChatInputPart{
{
Type: codersdk.ChatInputPartTypeText,
Text: "adding another disabled server is rejected",
},
},
MCPServerIDs: &[]uuid.UUID{cfg.ID, secondCfg.ID},
})
sdkErr := requireSDKError(t, err, http.StatusBadRequest)
require.Equal(t, "One or more MCP server IDs are invalid or disabled.", sdkErr.Message)
require.Equal(t, "Invalid IDs: "+secondCfg.ID.String(), sdkErr.Detail)
})
t.Run("MCPServerIDsThirdOrgRejected", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
client, db := newChatClientWithDatabase(t)
coderdtest.CreateFirstUser(t, client.Client)
_ = createChatModelConfig(t, client)
thirdOrg := dbgen.Organization(t, db, database.Organization{})
thirdOrgConfig := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{
OrganizationID: thirdOrg.ID,
Enabled: true,
})
secondOrg := dbgen.Organization(t, db, database.Organization{})
memberClientRaw, _ := coderdtest.CreateAnotherUser(t, client.Client, secondOrg.ID, rbac.ScopedRoleAgentsAccess(secondOrg.ID))
memberClient := codersdk.NewExperimentalClient(memberClientRaw)
_, err := memberClient.CreateChat(ctx, codersdk.CreateChatRequest{
OrganizationID: secondOrg.ID,
Content: []codersdk.ChatInputPart{
{
Type: codersdk.ChatInputPartTypeText,
Text: "chat with a third-org MCP server",
},
},
MCPServerIDs: []uuid.UUID{thirdOrgConfig.ID},
})
sdkErr := requireSDKError(t, err, http.StatusBadRequest)
require.Equal(t, "One or more MCP server IDs are invalid or disabled.", sdkErr.Message)
require.Equal(t, "Invalid IDs: "+thirdOrgConfig.ID.String(), sdkErr.Detail)
})
t.Run("MemberWithoutAgentsAccess", func(t *testing.T) {
t.Parallel()
@@ -1142,7 +1385,7 @@ func TestChats_ForceOnMCPServerEnforced(t *testing.T) {
_ = createChatModelConfig(t, client)
// An admin marks an MCP server as Force On.
forced, err := client.Client.CreateMCPServerConfig(ctx, codersdk.CreateMCPServerConfigRequest{
forced, err := client.Client.CreateMCPServerConfig(ctx, firstUser.OrganizationID, codersdk.CreateMCPServerConfigRequest{
DisplayName: "Forced Server",
Slug: "forced-server",
Transport: "streamable_http",
+57
View File
@@ -0,0 +1,57 @@
package httpmw
import (
"context"
"net/http"
"github.com/coder/coder/v2/coderd/database"
"github.com/coder/coder/v2/coderd/httpapi"
"github.com/coder/coder/v2/codersdk"
)
type mcpServerConfigParamContextKey struct{}
// MCPServerConfigParam returns the MCP server config from the
// ExtractMCPServerConfigParam handler.
func MCPServerConfigParam(r *http.Request) database.MCPServerConfig {
config, ok := r.Context().Value(mcpServerConfigParamContextKey{}).(database.MCPServerConfig)
if !ok {
panic("developer error: mcp server config param middleware not provided")
}
return config
}
// ExtractMCPServerConfigParam reads the "mcpserverconfig" URL parameter.
// Unauthorized reads are concealed as not found, so denied and missing rows
// both return 404.
func ExtractMCPServerConfigParam(db database.Store) func(http.Handler) http.Handler {
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(rw http.ResponseWriter, r *http.Request) {
ctx := r.Context()
configID, parsed := ParseUUIDParam(rw, r, "mcpserverconfig")
if !parsed {
return
}
config, err := db.GetMCPServerConfigByID(ctx, configID)
if httpapi.Is404Error(err) {
httpapi.ResourceNotFound(rw)
return
}
if err != nil {
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
Message: "Internal error fetching MCP server config.",
Detail: err.Error(),
})
return
}
if config.OrganizationID != OrganizationParam(r).ID {
httpapi.ResourceNotFound(rw)
return
}
ctx = context.WithValue(ctx, mcpServerConfigParamContextKey{}, config)
next.ServeHTTP(rw, r.WithContext(ctx))
})
}
}
+191 -242
View File
@@ -151,19 +151,24 @@ func shouldRefreshOIDCToken(link database.UserLink) (bool, time.Time) {
func (api *API) listMCPServerConfigs(rw http.ResponseWriter, r *http.Request) {
ctx := r.Context()
apiKey := httpmw.APIKey(r)
organization := httpmw.OrganizationParam(r)
// Admin users can see all MCP server configs (including disabled
// ones) for management purposes. Non-admin users see only enabled
// configs, which is sufficient for using the chat feature.
isAdmin := api.Authorize(r, policy.ActionRead, rbac.ResourceDeploymentConfig)
// Full view: disabled configs included, management fields unredacted.
// Auditors get it to inspect audit-logged resources; their MCP config
// read grant cannot select it because members hold the same read.
// Other members see enabled configs with management fields redacted.
// The update leg also requires config read so a custom role granting
// update without read cannot lift the read filtering below.
hasFullView := (api.Authorize(r, policy.ActionRead, rbac.ResourceMCPServerConfig.InOrg(organization.ID)) &&
api.Authorize(r, policy.ActionUpdate, rbac.ResourceMCPServerConfig.InOrg(organization.ID))) ||
api.Authorize(r, policy.ActionRead, rbac.ResourceAuditLog.InOrg(organization.ID))
var configs []database.MCPServerConfig
var err error
if isAdmin {
configs, err = api.Database.GetMCPServerConfigs(ctx)
if hasFullView {
configs, err = api.Database.GetMCPServerConfigsByOrganization(ctx, organization.ID)
} else {
//nolint:gocritic // All authenticated users need to read enabled MCP server configs to use the chat feature.
configs, err = api.Database.GetEnabledMCPServerConfigs(dbauthz.AsSystemRestricted(ctx))
configs, err = api.Database.GetEnabledMCPServerConfigsByOrganization(ctx, organization.ID)
}
if err != nil {
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
@@ -176,7 +181,7 @@ func (api *API) listMCPServerConfigs(rw http.ResponseWriter, r *http.Request) {
// Look up the calling user's OAuth2 tokens so we can populate
// auth_connected per server. Attempt to refresh expired tokens
// so the status is accurate and the token is ready for use.
//nolint:gocritic // Need to check user tokens across all servers.
//nolint:gocritic // Token authorization is handled separately from config RBAC.
userTokens, err := api.Database.GetMCPServerUserTokensByUserID(dbauthz.AsSystemRestricted(ctx), apiKey.UserID)
if err != nil {
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
@@ -204,7 +209,7 @@ func (api *API) listMCPServerConfigs(rw http.ResponseWriter, r *http.Request) {
resp := make([]codersdk.MCPServerConfig, 0, len(configs))
for _, config := range configs {
var sdkConfig codersdk.MCPServerConfig
if isAdmin {
if hasFullView {
sdkConfig = convertMCPServerConfig(config)
} else {
sdkConfig = convertMCPServerConfigRedacted(config)
@@ -226,7 +231,8 @@ func (api *API) listMCPServerConfigs(rw http.ResponseWriter, r *http.Request) {
func (api *API) createMCPServerConfig(rw http.ResponseWriter, r *http.Request) {
ctx := r.Context()
apiKey := httpmw.APIKey(r)
if !api.Authorize(r, policy.ActionUpdate, rbac.ResourceDeploymentConfig) {
organization := httpmw.OrganizationParam(r)
if !api.Authorize(r, policy.ActionCreate, rbac.ResourceMCPServerConfig.InOrg(organization.ID)) {
httpapi.Forbidden(rw)
return
}
@@ -236,6 +242,10 @@ func (api *API) createMCPServerConfig(rw http.ResponseWriter, r *http.Request) {
return
}
if req.AuthType == "user_oidc" && !api.authorizeUserOIDCMCPServerConfig(rw, r) {
return
}
if trimmed := strings.TrimSpace(req.OAuth2RevocationURL); trimmed != "" {
if err := mcpclient.ValidateRevocationEndpoint(trimmed); err != nil {
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
@@ -246,6 +256,8 @@ func (api *API) createMCPServerConfig(rw http.ResponseWriter, r *http.Request) {
}
}
configID := uuid.New()
// Validate auth-type-dependent fields.
switch req.AuthType {
case "oauth2":
@@ -256,74 +268,28 @@ func (api *API) createMCPServerConfig(rw http.ResponseWriter, r *http.Request) {
// Metadata (RFC 9728) and Authorization Server Metadata
// (RFC 8414), then register a client dynamically.
if req.OAuth2ClientID == "" && req.OAuth2AuthURL == "" && req.OAuth2TokenURL == "" {
// Auto-discovery flow: we need the config ID first to
// build the correct callback URL. Insert the record
// with empty OAuth2 fields, perform discovery, then
// update.
customHeadersJSON, err := marshalCustomHeaders(req.CustomHeaders)
if err != nil {
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
Message: "Invalid custom headers.",
// Create-only callers cannot read configs. This pre-DCR check reveals nothing
// beyond the insert's conflict response, which remains authoritative for races.
//nolint:gocritic // Restrict system access to this existence check.
_, err := api.Database.GetMCPServerConfigByOrganizationAndSlug(dbauthz.AsSystemRestricted(ctx), database.GetMCPServerConfigByOrganizationAndSlugParams{
OrganizationID: organization.ID,
Slug: strings.TrimSpace(req.Slug),
})
switch {
case err == nil:
httpapi.Write(ctx, rw, http.StatusConflict, codersdk.Response{
Message: "MCP server config already exists.",
})
return
case !errors.Is(err, sql.ErrNoRows):
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
Message: "Failed to create MCP server config.",
Detail: err.Error(),
})
return
}
inserted, err := api.Database.InsertMCPServerConfig(ctx, database.InsertMCPServerConfigParams{
DisplayName: strings.TrimSpace(req.DisplayName),
Slug: strings.TrimSpace(req.Slug),
Description: strings.TrimSpace(req.Description),
IconURL: strings.TrimSpace(req.IconURL),
Transport: strings.TrimSpace(req.Transport),
Url: strings.TrimSpace(req.URL),
AuthType: strings.TrimSpace(req.AuthType),
OAuth2ClientID: "",
OAuth2ClientSecret: "",
OAuth2ClientSecretKeyID: sql.NullString{},
OAuth2AuthURL: "",
OAuth2TokenURL: "",
OAuth2RevocationURL: "",
OAuth2Scopes: "",
APIKeyHeader: strings.TrimSpace(req.APIKeyHeader),
APIKeyValue: strings.TrimSpace(req.APIKeyValue),
APIKeyValueKeyID: sql.NullString{},
CustomHeaders: customHeadersJSON,
CustomHeadersKeyID: sql.NullString{},
ToolAllowList: coalesceStringSlice(trimStringSlice(req.ToolAllowList)),
ToolDenyList: coalesceStringSlice(trimStringSlice(req.ToolDenyList)),
Availability: strings.TrimSpace(req.Availability),
Enabled: req.Enabled,
ModelIntent: req.ModelIntent,
AllowInPlanMode: req.AllowInPlanMode,
ForwardCoderHeaders: req.ForwardCoderHeaders,
CreatedBy: apiKey.UserID,
UpdatedBy: apiKey.UserID,
})
if err != nil {
switch {
case database.IsUniqueViolation(err):
httpapi.Write(ctx, rw, http.StatusConflict, codersdk.Response{
Message: "MCP server config already exists.",
Detail: err.Error(),
})
return
case database.IsCheckViolation(err):
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
Message: "Invalid MCP server config.",
Detail: err.Error(),
})
return
default:
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
Message: "Failed to create MCP server config.",
Detail: err.Error(),
})
return
}
}
// Now build the callback URL with the actual ID.
callbackURL := fmt.Sprintf("%s/api/experimental/mcp/servers/%s/oauth2/callback", api.AccessURL.String(), inserted.ID)
callbackURL := api.AccessURL.String() + mcpServerOAuth2CallbackPath(configID)
// Discovery targets are attacker-influenced (the MCP
// server URL and any endpoints or redirects it
// advertises), so all discovery traffic goes through an
@@ -332,15 +298,6 @@ func (api *API) createMCPServerConfig(rw http.ResponseWriter, r *http.Request) {
httpClient := newMCPDiscoveryHTTPClient(api.HTTPClient, api.MCPOAuth2DiscoveryAllowedIPRanges)
result, err := discoverAndRegisterMCPOAuth2(ctx, httpClient, strings.TrimSpace(req.URL), callbackURL)
if err != nil {
// Clean up: delete the partially created config.
deleteErr := api.Database.DeleteMCPServerConfigByID(ctx, inserted.ID)
if deleteErr != nil {
api.Logger.Warn(ctx, "failed to clean up MCP server config after OAuth2 discovery failure",
slog.F("config_id", inserted.ID),
slog.Error(deleteErr),
)
}
api.Logger.Warn(ctx, "mcp oauth2 auto-discovery failed",
slog.F("url", req.URL),
slog.Error(err),
@@ -375,47 +332,12 @@ func (api *API) createMCPServerConfig(rw http.ResponseWriter, r *http.Request) {
}
}
// Update the record with discovered OAuth2 credentials.
updated, err := api.Database.UpdateMCPServerConfig(ctx, database.UpdateMCPServerConfigParams{
ID: inserted.ID,
DisplayName: inserted.DisplayName,
Slug: inserted.Slug,
Description: inserted.Description,
IconURL: inserted.IconURL,
Transport: inserted.Transport,
Url: inserted.Url,
AuthType: inserted.AuthType,
OAuth2ClientID: result.clientID,
OAuth2ClientSecret: result.clientSecret,
OAuth2ClientSecretKeyID: sql.NullString{},
OAuth2AuthURL: result.authURL,
OAuth2TokenURL: result.tokenURL,
OAuth2RevocationURL: oauth2RevocationURL,
OAuth2Scopes: oauth2Scopes,
APIKeyHeader: inserted.APIKeyHeader,
APIKeyValue: inserted.APIKeyValue,
APIKeyValueKeyID: inserted.APIKeyValueKeyID,
CustomHeaders: inserted.CustomHeaders,
CustomHeadersKeyID: inserted.CustomHeadersKeyID,
ToolAllowList: inserted.ToolAllowList,
ToolDenyList: inserted.ToolDenyList,
Availability: inserted.Availability,
Enabled: inserted.Enabled,
ModelIntent: inserted.ModelIntent,
AllowInPlanMode: inserted.AllowInPlanMode,
ForwardCoderHeaders: inserted.ForwardCoderHeaders,
UpdatedBy: apiKey.UserID,
})
if err != nil {
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
Message: "Failed to update MCP server config with OAuth2 credentials.",
Detail: err.Error(),
})
return
}
httpapi.Write(ctx, rw, http.StatusCreated, convertMCPServerConfig(updated))
return
req.OAuth2ClientID = result.clientID
req.OAuth2ClientSecret = result.clientSecret
req.OAuth2AuthURL = result.authURL
req.OAuth2TokenURL = result.tokenURL
req.OAuth2RevocationURL = oauth2RevocationURL
req.OAuth2Scopes = oauth2Scopes
} else if req.OAuth2ClientID == "" || req.OAuth2AuthURL == "" || req.OAuth2TokenURL == "" {
// Partial manual config: all three fields are required together.
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
@@ -449,6 +371,8 @@ func (api *API) createMCPServerConfig(rw http.ResponseWriter, r *http.Request) {
}
inserted, err := api.Database.InsertMCPServerConfig(ctx, database.InsertMCPServerConfigParams{
ID: configID,
OrganizationID: organization.ID,
DisplayName: strings.TrimSpace(req.DisplayName),
Slug: strings.TrimSpace(req.Slug),
Description: strings.TrimSpace(req.Description),
@@ -512,40 +436,18 @@ func (api *API) createMCPServerConfig(rw http.ResponseWriter, r *http.Request) {
func (api *API) getMCPServerConfig(rw http.ResponseWriter, r *http.Request) {
ctx := r.Context()
apiKey := httpmw.APIKey(r)
config := httpmw.MCPServerConfigParam(r)
mcpServerID, ok := parseMCPServerConfigID(rw, r)
if !ok {
return
}
isAdmin := api.Authorize(r, policy.ActionRead, rbac.ResourceDeploymentConfig)
var config database.MCPServerConfig
var err error
if isAdmin {
config, err = api.Database.GetMCPServerConfigByID(ctx, mcpServerID)
} else {
//nolint:gocritic // All authenticated users can view enabled MCP server configs.
config, err = api.Database.GetMCPServerConfigByID(dbauthz.AsSystemRestricted(ctx), mcpServerID)
if err == nil && !config.Enabled {
httpapi.ResourceNotFound(rw)
return
}
}
if err != nil {
if httpapi.Is404Error(err) {
httpapi.ResourceNotFound(rw)
return
}
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
Message: "Failed to get MCP server config.",
Detail: err.Error(),
})
// Same full-view rule as listMCPServerConfigs: admins and auditors.
hasFullView := api.Authorize(r, policy.ActionUpdate, config) ||
api.Authorize(r, policy.ActionRead, rbac.ResourceAuditLog.InOrg(config.OrganizationID))
if !hasFullView && !config.Enabled {
httpapi.ResourceNotFound(rw)
return
}
var sdkConfig codersdk.MCPServerConfig
if isAdmin {
if hasFullView {
sdkConfig = convertMCPServerConfig(config)
} else {
sdkConfig = convertMCPServerConfigRedacted(config)
@@ -554,26 +456,52 @@ func (api *API) getMCPServerConfig(rw http.ResponseWriter, r *http.Request) {
// Populate AuthConnected for the calling user. Attempt to
// refresh the token so the status is accurate.
if config.AuthType == "oauth2" {
//nolint:gocritic // Need to check user token for this server.
userTokens, err := api.Database.GetMCPServerUserTokensByUserID(dbauthz.AsSystemRestricted(ctx), apiKey.UserID)
if err != nil {
//nolint:gocritic // Token authorization is handled separately from config RBAC.
tok, err := api.Database.GetMCPServerUserToken(dbauthz.AsSystemRestricted(ctx), database.GetMCPServerUserTokenParams{
MCPServerConfigID: config.ID,
UserID: apiKey.UserID,
})
if err == nil {
sdkConfig.AuthConnected = api.refreshMCPUserToken(ctx, config, tok)
} else if !errors.Is(err, sql.ErrNoRows) {
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
Message: "Failed to get user tokens.",
Message: "Failed to get user token.",
Detail: err.Error(),
})
return
}
for _, tok := range userTokens {
if tok.MCPServerConfigID == config.ID {
sdkConfig.AuthConnected = api.refreshMCPUserToken(ctx, config, tok)
break
}
}
}
httpapi.Write(ctx, rw, http.StatusOK, sdkConfig)
}
var errUserOIDCRequiresDeploymentPerms = xerrors.New("managing user_oidc MCP server configs requires deployment-level permissions")
var errMCPConfigSupersededDuringAuth = xerrors.New("MCP server config superseded during authorization")
// authorizeUserOIDCMCPServerConfig requires deployment-level permission because
// user_oidc sends each chat owner's upstream OIDC access token to the configured
// URL without a per-user consent step.
func (api *API) authorizeUserOIDCMCPServerConfig(rw http.ResponseWriter, r *http.Request) bool {
if api.Authorize(r, policy.ActionUpdate, rbac.ResourceDeploymentConfig) {
return true
}
httpapi.Write(r.Context(), rw, http.StatusForbidden, codersdk.Response{
Message: "Managing user_oidc MCP server configs requires deployment-level permissions.",
})
return false
}
// Preserve the param middleware's 404 concealment. Write denial is a 403.
func (api *API) getMCPServerConfigForMutation(rw http.ResponseWriter, r *http.Request, action policy.Action) (database.MCPServerConfig, bool) {
config := httpmw.MCPServerConfigParam(r)
if !api.Authorize(r, action, config) {
httpapi.Forbidden(rw)
return database.MCPServerConfig{}, false
}
return config, true
}
// @Summary Update MCP server config
// @x-apidocgen {"skip": true}
// EXPERIMENTAL: this endpoint is experimental and is subject to change.
@@ -582,12 +510,7 @@ func (api *API) getMCPServerConfig(rw http.ResponseWriter, r *http.Request) {
func (api *API) updateMCPServerConfig(rw http.ResponseWriter, r *http.Request) {
ctx := r.Context()
apiKey := httpmw.APIKey(r)
if !api.Authorize(r, policy.ActionUpdate, rbac.ResourceDeploymentConfig) {
httpapi.Forbidden(rw)
return
}
mcpServerID, ok := parseMCPServerConfigID(rw, r)
existing, ok := api.getMCPServerConfigForMutation(rw, r, policy.ActionUpdate)
if !ok {
return
}
@@ -636,10 +559,20 @@ func (api *API) updateMCPServerConfig(rw http.ResponseWriter, r *http.Request) {
var updated database.MCPServerConfig
err := api.Database.InTx(func(tx database.Store) error {
existing, err := tx.GetMCPServerConfigByID(ctx, mcpServerID)
// Lock and re-fetch the row so omitted fields come from the latest
// version and grant invalidation serializes with in-flight OAuth
// callbacks verifying the same config.
current, err := tx.GetMCPServerConfigByIDForUpdate(ctx, existing.ID)
if err != nil {
return err
}
existing = current
touchesUserOIDC := existing.AuthType == "user_oidc" ||
(req.AuthType != nil && *req.AuthType == "user_oidc")
if touchesUserOIDC && !api.Authorize(r, policy.ActionUpdate, rbac.ResourceDeploymentConfig) {
return errUserOIDCRequiresDeploymentPerms
}
displayName := existing.DisplayName
if req.DisplayName != nil {
@@ -828,7 +761,18 @@ func (api *API) updateMCPServerConfig(rw http.ResponseWriter, r *http.Request) {
}
}
updated, err = tx.UpdateMCPServerConfig(ctx, database.UpdateMCPServerConfigParams{
// User grants are bound to the destination, auth flow, token and revocation
// endpoints, and OAuth client. Invalidate them when any of these change so
// stored tokens cannot be sent to another endpoint or client.
if serverURL != existing.Url || authType != existing.AuthType ||
oauth2TokenURL != existing.OAuth2TokenURL || oauth2RevocationURL != existing.OAuth2RevocationURL ||
oauth2ClientID != existing.OAuth2ClientID {
if err := tx.DeleteMCPServerUserTokensByConfigID(ctx, existing.ID); err != nil {
return xerrors.Errorf("invalidate MCP server user tokens: %w", err)
}
}
updatedConfig, err := tx.UpdateMCPServerConfig(ctx, database.UpdateMCPServerConfigParams{
DisplayName: displayName,
Slug: slug,
Description: description,
@@ -858,10 +802,19 @@ func (api *API) updateMCPServerConfig(rw http.ResponseWriter, r *http.Request) {
UpdatedBy: apiKey.UserID,
ID: existing.ID,
})
return err
if err != nil {
return err
}
updated = updatedConfig
return nil
}, nil)
if err != nil {
switch {
case errors.Is(err, errUserOIDCRequiresDeploymentPerms):
httpapi.Write(ctx, rw, http.StatusForbidden, codersdk.Response{
Message: "Managing user_oidc MCP server configs requires deployment-level permissions.",
})
return
case httpapi.Is404Error(err):
httpapi.ResourceNotFound(rw)
return
@@ -894,29 +847,12 @@ func (api *API) updateMCPServerConfig(rw http.ResponseWriter, r *http.Request) {
// EXPERIMENTAL: this endpoint is experimental and is subject to change.
func (api *API) deleteMCPServerConfig(rw http.ResponseWriter, r *http.Request) {
ctx := r.Context()
if !api.Authorize(r, policy.ActionUpdate, rbac.ResourceDeploymentConfig) {
httpapi.Forbidden(rw)
return
}
mcpServerID, ok := parseMCPServerConfigID(rw, r)
config, ok := api.getMCPServerConfigForMutation(rw, r, policy.ActionDelete)
if !ok {
return
}
if _, err := api.Database.GetMCPServerConfigByID(ctx, mcpServerID); err != nil {
if httpapi.Is404Error(err) {
httpapi.ResourceNotFound(rw)
return
}
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
Message: "Failed to get MCP server config.",
Detail: err.Error(),
})
return
}
if err := api.Database.DeleteMCPServerConfigByID(ctx, mcpServerID); err != nil {
if err := api.Database.DeleteMCPServerConfigByID(ctx, config.ID); err != nil {
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
Message: "Failed to delete MCP server config.",
Detail: err.Error(),
@@ -935,25 +871,7 @@ func (api *API) deleteMCPServerConfig(rw http.ResponseWriter, r *http.Request) {
//nolint:revive // HTTP handler writes to ResponseWriter.
func (api *API) mcpServerOAuth2Connect(rw http.ResponseWriter, r *http.Request) {
ctx := r.Context()
mcpServerID, ok := parseMCPServerConfigID(rw, r)
if !ok {
return
}
//nolint:gocritic // Any authenticated user can initiate OAuth2 for an enabled MCP server.
config, err := api.Database.GetMCPServerConfigByID(dbauthz.AsSystemRestricted(ctx), mcpServerID)
if err != nil {
if httpapi.Is404Error(err) {
httpapi.ResourceNotFound(rw)
return
}
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
Message: "Failed to get MCP server config.",
Detail: err.Error(),
})
return
}
config := httpmw.MCPServerConfigParam(r)
if !config.Enabled {
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
@@ -980,7 +898,7 @@ func (api *API) mcpServerOAuth2Connect(rw http.ResponseWriter, r *http.Request)
// The callback URL is on our server; after the exchange we store
// the token and close the popup.
state := uuid.New().String()
callbackPath := fmt.Sprintf("/api/experimental/mcp/servers/%s/oauth2/callback", config.ID)
callbackPath := mcpServerOAuth2CallbackPath(config.ID)
http.SetCookie(rw, api.DeploymentValues.HTTPCookies.Apply(&http.Cookie{
Name: "mcp_oauth2_state_" + config.ID.String(),
Value: state,
@@ -1035,18 +953,13 @@ func (api *API) mcpServerOAuth2Callback(rw http.ResponseWriter, r *http.Request)
if !ok {
return
}
//nolint:gocritic // Any authenticated user can complete OAuth2 for an enabled MCP server.
config, err := api.Database.GetMCPServerConfigByID(dbauthz.AsSystemRestricted(ctx), mcpServerID)
config, err := api.Database.GetMCPServerConfigByID(ctx, mcpServerID)
if err != nil {
if httpapi.Is404Error(err) {
httpapi.ResourceNotFound(rw)
return
}
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
Message: "Failed to get MCP server config.",
Detail: err.Error(),
})
httpapi.InternalServerError(rw, err)
return
}
@@ -1099,7 +1012,7 @@ func (api *API) mcpServerOAuth2Callback(rw http.ResponseWriter, r *http.Request)
return
}
// Clear the state cookie.
callbackPath := fmt.Sprintf("/api/experimental/mcp/servers/%s/oauth2/callback", config.ID)
callbackPath := mcpServerOAuth2CallbackPath(config.ID)
http.SetCookie(rw, api.DeploymentValues.HTTPCookies.Apply(&http.Cookie{
Name: "mcp_oauth2_state_" + config.ID.String(),
Value: "",
@@ -1168,17 +1081,39 @@ func (api *API) mcpServerOAuth2Callback(rw http.ResponseWriter, r *http.Request)
expiry = sql.NullTime{Time: token.Expiry, Valid: true}
}
//nolint:gocritic // Users store their own tokens.
_, err = api.Database.UpsertMCPServerUserToken(dbauthz.AsSystemRestricted(ctx), database.UpsertMCPServerUserTokenParams{
MCPServerConfigID: mcpServerID,
UserID: apiKey.UserID,
AccessToken: token.AccessToken,
AccessTokenKeyID: sql.NullString{},
RefreshToken: refreshToken,
RefreshTokenKeyID: sql.NullString{},
TokenType: token.TokenType,
Expiry: expiry,
})
err = api.Database.InTx(func(tx database.Store) error {
// Hold the config lock through the grant write so a concurrent update
// cannot invalidate grants and then have this callback recreate one
// for the old config.
current, err := tx.GetMCPServerConfigByIDForUpdate(ctx, config.ID)
if err != nil {
return xerrors.Errorf("re-read MCP server config: %w", err)
}
if current.Url != config.Url || current.AuthType != config.AuthType ||
current.OAuth2TokenURL != config.OAuth2TokenURL || current.OAuth2RevocationURL != config.OAuth2RevocationURL ||
current.OAuth2ClientID != config.OAuth2ClientID {
return errMCPConfigSupersededDuringAuth
}
//nolint:gocritic // Users store their own tokens.
_, err = tx.UpsertMCPServerUserToken(dbauthz.AsSystemRestricted(ctx), database.UpsertMCPServerUserTokenParams{
MCPServerConfigID: config.ID,
UserID: apiKey.UserID,
AccessToken: token.AccessToken,
AccessTokenKeyID: sql.NullString{},
RefreshToken: refreshToken,
RefreshTokenKeyID: sql.NullString{},
TokenType: token.TokenType,
Expiry: expiry,
})
return err
}, nil)
if errors.Is(err, errMCPConfigSupersededDuringAuth) {
httpapi.Write(ctx, rw, http.StatusConflict, codersdk.Response{
Message: "MCP server configuration changed during authorization.",
Detail: "The server's destination or OAuth settings were updated while the connection was in progress. Reconnect to authorize against the current configuration.",
})
return
}
if err != nil {
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
Message: "Failed to store OAuth2 token.",
@@ -1210,8 +1145,8 @@ func (api *API) mcpServerOAuth2Disconnect(rw http.ResponseWriter, r *http.Reques
ctx := r.Context()
apiKey := httpmw.APIKey(r)
mcpServerID, ok := parseMCPServerConfigID(rw, r)
if !ok {
configID, parsed := httpmw.ParseUUIDParam(rw, r, "mcpServer")
if !parsed {
return
}
@@ -1224,20 +1159,23 @@ func (api *API) mcpServerOAuth2Disconnect(rw http.ResponseWriter, r *http.Reques
// 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,
MCPServerConfigID: configID,
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)
// without a token cannot probe which config IDs exist. The
// system context keeps disconnect available to token owners
// who can no longer read the config, such as users removed
// from the organization.
dbConfig, err := tx.GetMCPServerConfigByID(systemCtx, configID)
if err != nil {
return err
}
if err := tx.DeleteMCPServerUserToken(systemCtx, database.DeleteMCPServerUserTokenParams{
MCPServerConfigID: mcpServerID,
MCPServerConfigID: configID,
UserID: apiKey.UserID,
}); err != nil {
return err
@@ -1418,8 +1356,18 @@ func (api *API) markMCPTokenRefreshFailure(
return false
}
// mcpServerOAuth2CallbackPath returns the OAuth2 callback path for a
// config. This path is frozen: it is the redirect URI registered with
// external authorization servers, so it must not change when other MCP
// routes move. The route registration in coderd.go and the OAuth cookie
// Path values must stay aligned with it.
func mcpServerOAuth2CallbackPath(configID uuid.UUID) string {
return fmt.Sprintf("/api/experimental/mcp/servers/%s/oauth2/callback", configID)
}
// parseMCPServerConfigID extracts the MCP server config UUID from the
// "mcpServer" path parameter.
// "mcpServer" path parameter, which is part of the frozen callback
// route shape.
func parseMCPServerConfigID(rw http.ResponseWriter, r *http.Request) (uuid.UUID, bool) {
mcpServerID, err := uuid.Parse(chi.URLParam(r, "mcpServer"))
if err != nil {
@@ -1437,11 +1385,12 @@ func parseMCPServerConfigID(rw http.ResponseWriter, r *http.Request) (uuid.UUID,
// Admin-only fields (OAuth2 client ID, auth URLs, etc.) are included.
func convertMCPServerConfig(config database.MCPServerConfig) codersdk.MCPServerConfig {
return codersdk.MCPServerConfig{
ID: config.ID,
DisplayName: config.DisplayName,
Slug: config.Slug,
Description: config.Description,
IconURL: config.IconURL,
ID: config.ID,
OrganizationID: config.OrganizationID,
DisplayName: config.DisplayName,
Slug: config.Slug,
Description: config.Description,
IconURL: config.IconURL,
Transport: config.Transport,
URL: config.Url,
+644 -143
View File
File diff suppressed because it is too large Load Diff
+11
View File
@@ -219,6 +219,16 @@ var (
Type: "license",
}
// ResourceMCPServerConfig
// Valid Actions
// - "ActionCreate" :: create a new MCP server config
// - "ActionDelete" :: delete MCP server config
// - "ActionRead" :: read MCP server config
// - "ActionUpdate" :: update MCP server config
ResourceMCPServerConfig = Object{
Type: "mcp_server_config",
}
// ResourceNotificationMessage
// Valid Actions
// - "ActionCreate" :: create notification messages
@@ -522,6 +532,7 @@ func AllResources() []Objecter {
ResourceIdpsyncSettings,
ResourceInboxNotification,
ResourceLicense,
ResourceMCPServerConfig,
ResourceNotificationMessage,
ResourceNotificationPreference,
ResourceNotificationTemplate,
+11
View File
@@ -85,6 +85,13 @@ var chatActions = map[Action]ActionDefinition{
ActionShare: "share a chat with other users or groups",
}
var mcpServerConfigActions = map[Action]ActionDefinition{
ActionCreate: "create a new MCP server config",
ActionRead: "read MCP server config",
ActionUpdate: "update MCP server config",
ActionDelete: "delete MCP server config",
}
// RBACPermissions is indexed by the type
var RBACPermissions = map[string]PermissionDefinition{
// Wildcard is every object, and the action "*" provides all actions.
@@ -453,4 +460,8 @@ var RBACPermissions = map[string]PermissionDefinition{
ActionDelete: "delete boundary usage statistics",
},
},
"mcp_server_config": {
Name: "MCPServerConfig",
Actions: mcpServerConfigActions,
},
}
+16
View File
@@ -74,6 +74,22 @@ func ChatNoACLConverter() *sqltypes.VariableConverter {
return matcher
}
// MCPServerConfigNoACLConverter converts MCP server config permissions to SQL.
// Until sharing adds ACL columns, ACL matchers stay false and only organization
// ownership filters rows.
func MCPServerConfigNoACLConverter() *sqltypes.VariableConverter {
matcher := sqltypes.NewVariableConverter().RegisterMatcher(
resourceIDMatcher(),
organizationOwnerMatcher(),
sqltypes.AlwaysFalse(userOwnerMatcher()),
)
matcher.RegisterMatcher(
sqltypes.AlwaysFalse(groupACLMatcher(matcher)),
sqltypes.AlwaysFalse(userACLMatcher(matcher)),
)
return matcher
}
func chatBaseConverter() *sqltypes.VariableConverter {
return sqltypes.NewVariableConverter().RegisterMatcher(
chatResourceIDMatcher(),
+8
View File
@@ -476,6 +476,7 @@ func ReloadBuiltinRoles(opts *RoleOptions) {
// Allow auditors to query deployment stats and insights.
ResourceDeploymentStats.Type: {policy.ActionRead},
ResourceDeploymentConfig.Type: {policy.ActionRead},
ResourceMCPServerConfig.Type: {policy.ActionRead},
// Allow auditors to query AI Bridge interceptions.
ResourceAibridgeInterception.Type: {policy.ActionRead},
// Allow auditors to read boundary logs.
@@ -611,6 +612,7 @@ func ReloadBuiltinRoles(opts *RoleOptions) {
ResourceGroupMember.Type: {policy.ActionRead},
ResourceOrganization.Type: {policy.ActionRead},
ResourceOrganizationMember.Type: {policy.ActionRead},
ResourceMCPServerConfig.Type: {policy.ActionRead},
}),
Member: []Permission{},
},
@@ -1157,6 +1159,9 @@ func OrgMemberPermissions(org OrgSettings) OrgRolePermissions {
ResourceOrganization.Type: {policy.ActionRead},
// Can read available roles.
ResourceAssignOrgRole.Type: {policy.ActionRead},
// TODO(mafredri): Remove once CODAGT-712 replaces this grant with
// per-config ACL evaluation.
ResourceMCPServerConfig.Type: {policy.ActionRead},
}
// In all modes of workspace sharing but `none`, members need to
@@ -1234,6 +1239,9 @@ func OrgServiceAccountPermissions(org OrgSettings) OrgRolePermissions {
ResourceOrganization.Type: {policy.ActionRead},
// Can read available roles.
ResourceAssignOrgRole.Type: {policy.ActionRead},
// TODO(mafredri): Remove once CODAGT-712 replaces this grant with
// per-config ACL evaluation.
ResourceMCPServerConfig.Type: {policy.ActionRead},
}
// When workspace sharing is enabled, service accounts need to see
+18
View File
@@ -824,6 +824,24 @@ func TestRolePermissions(t *testing.T) {
false: {setOtherOrg, setOrgNotMe, memberMe, agentsAccessUser, templateAdmin, userAdmin, orgWorkspaceAccessUser},
},
},
{
Name: "MCPServerConfigRead",
Actions: []policy.Action{policy.ActionRead},
Resource: rbac.ResourceMCPServerConfig.WithID(uuid.New()).InOrg(orgID),
AuthorizeMap: map[bool][]hasAuthSubjects{
true: {owner, orgAdmin, orgAuditor, auditor, orgMemberMe},
false: {setOtherOrg, memberMe, agentsAccessUser, orgWorkspaceAccessUser, orgUserAdmin, orgTemplateAdmin, templateAdmin, userAdmin},
},
},
{
Name: "MCPServerConfigWrite",
Actions: []policy.Action{policy.ActionCreate, policy.ActionUpdate, policy.ActionDelete},
Resource: rbac.ResourceMCPServerConfig.WithID(uuid.New()).InOrg(orgID),
AuthorizeMap: map[bool][]hasAuthSubjects{
true: {owner, orgAdmin},
false: {setOtherOrg, orgAuditor, auditor, orgMemberMe, memberMe, agentsAccessUser, orgWorkspaceAccessUser, orgUserAdmin, orgTemplateAdmin, templateAdmin, userAdmin},
},
},
{
Name: "DebugInfo",
Actions: []policy.Action{policy.ActionRead},
+12
View File
@@ -73,6 +73,10 @@ const (
ScopeLicenseCreate ScopeName = "license:create"
ScopeLicenseDelete ScopeName = "license:delete"
ScopeLicenseRead ScopeName = "license:read"
ScopeMcpServerConfigCreate ScopeName = "mcp_server_config:create"
ScopeMcpServerConfigDelete ScopeName = "mcp_server_config:delete"
ScopeMcpServerConfigRead ScopeName = "mcp_server_config:read"
ScopeMcpServerConfigUpdate ScopeName = "mcp_server_config:update"
ScopeNotificationMessageCreate ScopeName = "notification_message:create"
ScopeNotificationMessageDelete ScopeName = "notification_message:delete"
ScopeNotificationMessageRead ScopeName = "notification_message:read"
@@ -261,6 +265,10 @@ func (e ScopeName) Valid() bool {
ScopeLicenseCreate,
ScopeLicenseDelete,
ScopeLicenseRead,
ScopeMcpServerConfigCreate,
ScopeMcpServerConfigDelete,
ScopeMcpServerConfigRead,
ScopeMcpServerConfigUpdate,
ScopeNotificationMessageCreate,
ScopeNotificationMessageDelete,
ScopeNotificationMessageRead,
@@ -450,6 +458,10 @@ func AllScopeNameValues() []ScopeName {
ScopeLicenseCreate,
ScopeLicenseDelete,
ScopeLicenseRead,
ScopeMcpServerConfigCreate,
ScopeMcpServerConfigDelete,
ScopeMcpServerConfigRead,
ScopeMcpServerConfigUpdate,
ScopeNotificationMessageCreate,
ScopeNotificationMessageDelete,
ScopeNotificationMessageRead,
+9 -9
View File
@@ -1204,14 +1204,14 @@ type PromoteQueuedResult struct {
}
// enforceForcedMCPServerIDs appends the ID of every enabled Force On
// MCP server config missing from ids. Force On availability is a
// server-side policy: callers must not be able to exclude such
// servers by stripping IDs from a request (Cure53 CDM-02-010). The
// forced set is read with daemon scope because regular users cannot
// read MCP server configs directly.
func enforceForcedMCPServerIDs(ctx context.Context, store database.Store, ids []uuid.UUID) ([]uuid.UUID, error) {
// MCP server config in the chat's organization missing from ids. Force
// On availability is a server-side policy: callers must not be able to
// exclude such servers by stripping IDs from a request (Cure53
// CDM-02-010). The forced set is read with daemon scope because
// regular users cannot read MCP server configs directly.
func enforceForcedMCPServerIDs(ctx context.Context, store database.Store, organizationID uuid.UUID, ids []uuid.UUID) ([]uuid.UUID, error) {
//nolint:gocritic // Non-admin users need chatd-scoped config reads here.
forced, err := store.GetForcedMCPServerConfigs(dbauthz.AsChatd(ctx))
forced, err := store.GetForcedMCPServerConfigsByOrganization(dbauthz.AsChatd(ctx), organizationID)
if err != nil {
// Fail closed: proceeding without the forced set would
// silently bypass a security policy.
@@ -1258,7 +1258,7 @@ func (p *Server) CreateChat(ctx context.Context, opts CreateOptions) (database.C
// Force On MCP servers are enforced server-side so a caller
// cannot exclude them by stripping IDs from the request
// (Cure53 CDM-02-010).
enforcedMCPServerIDs, err := enforceForcedMCPServerIDs(ctx, p.db, opts.MCPServerIDs)
enforcedMCPServerIDs, err := enforceForcedMCPServerIDs(ctx, p.db, opts.OrganizationID, opts.MCPServerIDs)
if err != nil {
return database.Chat{}, err
}
@@ -1516,7 +1516,7 @@ func (p *Server) SendMessage(
// Force On MCP servers are enforced server-side so a
// caller cannot remove them by tampering with the
// update (Cure53 CDM-02-010).
enforcedIDs, enforceErr := enforceForcedMCPServerIDs(ctx, store, *requestedMCPServerIDs)
enforcedIDs, enforceErr := enforceForcedMCPServerIDs(ctx, store, lockedChat.OrganizationID, *requestedMCPServerIDs)
if enforceErr != nil {
return enforceErr
}
+85 -66
View File
@@ -405,7 +405,7 @@ func TestPlanModeSubagentChatExcludesAskUserQuestion(t *testing.T) {
mcpTS := httptest.NewServer(testMCPHTTPHandler(mcpSrv))
t.Cleanup(mcpTS.Close)
mcpConfig, err := client.CreateMCPServerConfig(ctx, codersdk.CreateMCPServerConfigRequest{
mcpConfig, err := client.CreateMCPServerConfig(ctx, user.OrganizationID, codersdk.CreateMCPServerConfigRequest{
DisplayName: "Plan Root MCP",
Slug: "plan-root-mcp",
Transport: "streamable_http",
@@ -732,18 +732,20 @@ func TestExploreChatUsesPersistedMCPSnapshot(t *testing.T) {
},
)
mcpConfig := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{
DisplayName: "External Snapshot MCP",
Slug: "external-snapshot-mcp",
Url: externalMCPServer.URL,
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
OrganizationID: org.ID,
DisplayName: "External Snapshot MCP",
Slug: "external-snapshot-mcp",
Url: externalMCPServer.URL,
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
})
dbgen.MCPServerConfig(t, db, database.MCPServerConfig{
DisplayName: "Second MCP",
Slug: "second-mcp",
Url: secondMCPServer.URL,
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
OrganizationID: org.ID,
DisplayName: "Second MCP",
Slug: "second-mcp",
Url: secondMCPServer.URL,
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
})
ws, dbAgent := seedWorkspaceWithAgent(t, db, user.ID)
@@ -866,11 +868,12 @@ func TestRootExploreChatStaysBuiltinOnlyAtRuntime(t *testing.T) {
user, org, model := seedChatDependenciesWithProvider(t, db, "openai-compat", openAIURL)
mcpConfig := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{
DisplayName: "Root Explore Runtime MCP",
Slug: "root-explore-runtime-mcp",
Url: externalMCPServer.URL,
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
OrganizationID: org.ID,
DisplayName: "Root Explore Runtime MCP",
Slug: "root-explore-runtime-mcp",
Url: externalMCPServer.URL,
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
})
factory := chattest.NewMockAIBridgeTransport(t, openAIURL)
@@ -1061,18 +1064,20 @@ func TestExploreChatSendMessageCannotMutateMCPSnapshot(t *testing.T) {
user, org, model := seedChatDependenciesWithProvider(t, db, "openai-compat", openAIURL)
parentConfig := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{
DisplayName: "Runtime Parent MCP",
Slug: "runtime-parent-mcp",
Url: parentTS.URL,
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
OrganizationID: org.ID,
DisplayName: "Runtime Parent MCP",
Slug: "runtime-parent-mcp",
Url: parentTS.URL,
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
})
injectedConfig := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{
DisplayName: "Runtime Injected MCP",
Slug: "runtime-injected-mcp",
Url: injectedTS.URL,
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
OrganizationID: org.ID,
DisplayName: "Runtime Injected MCP",
Slug: "runtime-injected-mcp",
Url: injectedTS.URL,
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
})
factory := chattest.NewMockAIBridgeTransport(t, openAIURL)
@@ -1191,6 +1196,7 @@ func TestPlanModeRootChatAllowsApprovedExternalMCPTools(t *testing.T) {
user, org, model := seedChatDependenciesWithProvider(t, db, "openai-compat", openAIURL)
approvedConfig := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{
OrganizationID: org.ID,
DisplayName: "Plan Approved MCP",
Slug: "plan-approved-mcp",
Url: echoTS.URL,
@@ -1200,14 +1206,16 @@ func TestPlanModeRootChatAllowsApprovedExternalMCPTools(t *testing.T) {
})
blockedConfig := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{
DisplayName: "Plan Blocked MCP",
Slug: "plan-blocked-mcp",
Url: echoTS.URL,
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
OrganizationID: org.ID,
DisplayName: "Plan Blocked MCP",
Slug: "plan-blocked-mcp",
Url: echoTS.URL,
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
})
filteredConfig := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{
OrganizationID: org.ID,
DisplayName: "Plan Filtered MCP",
Slug: "plan-filtered-mcp",
Url: filteredTS.URL,
@@ -10459,11 +10467,12 @@ func TestMCPToolSearchGenerationFlows(t *testing.T) {
})
user, org, model := seedChatDependenciesWithProvider(t, db, "openai-compat", openAIURL)
mcpConfig := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{
DisplayName: "Search MCP",
Slug: "search-mcp",
Url: mcpTS.URL,
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
OrganizationID: org.ID,
DisplayName: "Search MCP",
Slug: "search-mcp",
Url: mcpTS.URL,
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
})
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL))
@@ -10541,11 +10550,12 @@ func TestMCPToolSearchGenerationFlows(t *testing.T) {
})
user, org, model := seedChatDependenciesWithProvider(t, db, "openai-compat", openAIURL)
mcpConfig := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{
DisplayName: "Direct MCP",
Slug: "direct-mcp",
Url: mcpTS.URL,
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
OrganizationID: org.ID,
DisplayName: "Direct MCP",
Slug: "direct-mcp",
Url: mcpTS.URL,
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
})
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL))
@@ -10601,11 +10611,12 @@ func TestMCPToolSearchGenerationFlows(t *testing.T) {
})
user, org, model := seedChatDependenciesWithProvider(t, db, "openai-compat", openAIURL)
mcpConfig := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{
DisplayName: "Count MCP",
Slug: "count-mcp",
Url: mcpTS.URL,
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
OrganizationID: org.ID,
DisplayName: "Count MCP",
Slug: "count-mcp",
Url: mcpTS.URL,
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
})
reg := prometheus.NewRegistry()
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
@@ -10658,11 +10669,12 @@ func TestMCPToolSearchGenerationFlows(t *testing.T) {
})
user, org, model := seedChatDependenciesWithProvider(t, db, "openai-compat", openAIURL)
mcpConfig := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{
DisplayName: "Hooked MCP",
Slug: "hooked-mcp",
Url: mcpTS.URL,
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
OrganizationID: org.ID,
DisplayName: "Hooked MCP",
Slug: "hooked-mcp",
Url: mcpTS.URL,
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
})
reg := prometheus.NewRegistry()
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
@@ -10719,11 +10731,12 @@ func TestMCPToolSearchGenerationFlows(t *testing.T) {
t.Cleanup(consumer.Close)
user, org, model := seedChatDependenciesWithProvider(t, db, "openai-compat", openAIURL)
mcpConfig := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{
DisplayName: "Failing MCP",
Slug: "failing-mcp",
Url: mcpTS.URL,
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
OrganizationID: org.ID,
DisplayName: "Failing MCP",
Slug: "failing-mcp",
Url: mcpTS.URL,
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
})
reg := prometheus.NewRegistry()
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
@@ -10775,11 +10788,12 @@ func TestMCPToolSearchGenerationFlows(t *testing.T) {
model.ContextLimit = 100_000
model = updateChatModelContextLimit(t, db, model)
mcpConfig := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{
DisplayName: "Small MCP",
Slug: "small-mcp",
Url: mcpTS.URL,
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
OrganizationID: org.ID,
DisplayName: "Small MCP",
Slug: "small-mcp",
Url: mcpTS.URL,
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
})
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL))
@@ -10903,11 +10917,12 @@ func TestMCPServerToolInvocation(t *testing.T) {
// happen after seedChatDependencies so user.ID exists for
// the foreign key.
mcpConfig := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{
DisplayName: "Test MCP",
Slug: "test-mcp",
Url: mcpTS.URL,
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
OrganizationID: org.ID,
DisplayName: "Test MCP",
Slug: "test-mcp",
Url: mcpTS.URL,
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
})
ws, dbAgent := seedWorkspaceWithAgent(t, db, user.ID)
@@ -11065,6 +11080,7 @@ func TestPlanModeRootChatApprovedExternalMCPToolInvocation(t *testing.T) {
user, org, model := seedChatDependenciesWithProvider(t, db, "openai-compat", openAIURL)
mcpConfig := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{
OrganizationID: org.ID,
DisplayName: "Plan Mode MCP",
Slug: "plan-mode-mcp",
Url: mcpTS.URL,
@@ -11164,6 +11180,7 @@ func TestPlanModeRootChatApprovedExternalMCPWorkflowCanReachProposePlan(t *testi
user, org, model := seedChatDependenciesWithProvider(t, db, "openai-compat", openAIURL)
mcpConfig := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{
OrganizationID: org.ID,
DisplayName: "Plan Workflow MCP",
Slug: "plan-workflow-mcp",
Url: mcpTS.URL,
@@ -11364,6 +11381,7 @@ func TestMCPServerOAuth2TokenRefresh(t *testing.T) {
// Seed the MCP server config with OAuth2 auth pointing to our
// mock token endpoint.
mcpConfig := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{
OrganizationID: org.ID,
DisplayName: "Authed MCP",
Slug: "authed-mcp",
Url: mcpTS.URL,
@@ -11492,6 +11510,7 @@ func TestMCPServerOAuth2TokenRefreshFailureGraceful(t *testing.T) {
user, org, model := seedChatDependenciesWithProvider(t, db, "openai-compat", openAIURL)
mcpConfig := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{
OrganizationID: org.ID,
DisplayName: "Broken MCP",
Slug: "broken-mcp",
Url: "http://127.0.0.1:0/does-not-exist",
+41 -23
View File
@@ -81,12 +81,13 @@ func TestCreateChat_ForceOnMCPServerEnforced(t *testing.T) {
user, org, model := seedChatDependenciesWithProvider(t, db, "openai-compat", openAIURL)
forcedConfig := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{
DisplayName: "Forced MCP",
Slug: "forced-mcp",
Url: forcedURL,
Availability: "force_on",
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
OrganizationID: org.ID,
DisplayName: "Forced MCP",
Slug: "forced-mcp",
Url: forcedURL,
Availability: "force_on",
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
})
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
@@ -144,19 +145,21 @@ func TestSendMessage_ForceOnMCPServerEnforced(t *testing.T) {
user, org, model := seedChatDependenciesWithProvider(t, db, "openai-compat", openAIURL)
forcedConfig := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{
DisplayName: "Forced MCP",
Slug: "forced-mcp",
Url: forcedURL,
Availability: "force_on",
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
OrganizationID: org.ID,
DisplayName: "Forced MCP",
Slug: "forced-mcp",
Url: forcedURL,
Availability: "force_on",
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
})
optionalConfig := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{
DisplayName: "Optional MCP",
Slug: "optional-mcp",
Url: optionalURL,
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
OrganizationID: org.ID,
DisplayName: "Optional MCP",
Slug: "optional-mcp",
Url: optionalURL,
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
})
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
@@ -254,12 +257,25 @@ func TestGeneration_ForceOnMCPServerEnforcedForExistingChats(t *testing.T) {
// An admin marks a server force_on after the chat already exists.
dbgen.MCPServerConfig(t, db, database.MCPServerConfig{
DisplayName: "Forced MCP",
Slug: "forced-mcp",
Url: forcedURL,
Availability: "force_on",
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
OrganizationID: org.ID,
DisplayName: "Forced MCP",
Slug: "forced-mcp",
Url: forcedURL,
Availability: "force_on",
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
})
// A force_on server in another organization must not attach: the
// forced set is scoped to the chat's organization.
dbgen.MCPServerConfig(t, db, database.MCPServerConfig{
OrganizationID: dbgen.Organization(t, db, database.Organization{}).ID,
DisplayName: "Foreign Forced MCP",
Slug: "foreign-forced-mcp",
Url: forcedURL,
Availability: "force_on",
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
})
// A send that does not touch mcp_server_ids must still pick up
@@ -287,4 +303,6 @@ func TestGeneration_ForceOnMCPServerEnforcedForExistingChats(t *testing.T) {
"no force_on server existed during the first turn")
require.Contains(t, calls[len(calls)-1], "forced-mcp__echo",
"force_on MCP tools must reach generation for chats created before the policy")
require.NotContains(t, calls[len(calls)-1], "foreign-forced-mcp__echo",
"another organization's force_on server must not attach")
}
+18 -2
View File
@@ -42,7 +42,7 @@ func (server *Server) effectiveMCPServerConfigs(
var configs []database.MCPServerConfig
if len(chat.MCPServerIDs) > 0 {
var err error
configs, err = server.db.GetMCPServerConfigsByIDs(ctx, chat.MCPServerIDs)
configs, err = enabledMCPServerConfigsForChatOrg(ctx, server.db, chat.OrganizationID, chat.MCPServerIDs)
if err != nil {
// Best-effort for the user-selected set, matching prior
// behavior: a load failure degrades the turn rather than
@@ -54,7 +54,7 @@ func (server *Server) effectiveMCPServerConfigs(
if isExploreSubagentMode(chat.Mode) {
return configs, nil
}
forced, err := server.db.GetForcedMCPServerConfigs(ctx)
forced, err := server.db.GetForcedMCPServerConfigsByOrganization(ctx, chat.OrganizationID)
if err != nil {
// Fail closed: running the turn without the forced set would
// silently bypass a security policy.
@@ -910,3 +910,19 @@ func latestAssistantText(messages []database.ChatMessage) string {
}
return ""
}
func enabledMCPServerConfigsForChatOrg(
ctx context.Context,
db database.Store,
organizationID uuid.UUID,
ids []uuid.UUID,
) ([]database.MCPServerConfig, error) {
configs, err := db.GetEnabledMCPServerConfigsByOrganizationAndIDs(ctx, database.GetEnabledMCPServerConfigsByOrganizationAndIDsParams{
OrganizationID: organizationID,
IDs: ids,
})
if err != nil {
return nil, xerrors.Errorf("get enabled MCP server configs for organization: %w", err)
}
return configs, nil
}
@@ -21,6 +21,7 @@ import (
"github.com/coder/coder/v2/coderd/x/chatd/chatprovider"
"github.com/coder/coder/v2/coderd/x/chatd/chatstate"
"github.com/coder/coder/v2/codersdk"
"github.com/coder/coder/v2/testutil"
)
func mustMarshalText(t *testing.T, parts ...string) pqtype.NullRawMessage {
@@ -621,3 +622,130 @@ func TestShouldCompactPromptUsage(t *testing.T) {
contextLimit, 80))
})
}
func TestEnabledMCPServerConfigsForChatOrg(t *testing.T) {
t.Parallel()
newOrgWithConfig := func(t *testing.T, db database.Store) (database.Organization, database.MCPServerConfig) {
t.Helper()
org := dbgen.Organization(t, db, database.Organization{})
cfg := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{
OrganizationID: org.ID,
Enabled: true,
})
return org, cfg
}
t.Run("DefaultOrgConfigExcluded", func(t *testing.T) {
t.Parallel()
db, _ := dbtestutil.NewDB(t)
ctx := testutil.Context(t, testutil.WaitShort)
defaultOrg, err := db.GetDefaultOrganization(ctx)
require.NoError(t, err)
// Configs resolve only within the chat's organization.
chatOrg, chatOrgCfg := newOrgWithConfig(t, db)
defaultOrgCfg := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{
OrganizationID: defaultOrg.ID,
Enabled: true,
})
configs, err := enabledMCPServerConfigsForChatOrg(ctx, db, chatOrg.ID, []uuid.UUID{chatOrgCfg.ID, defaultOrgCfg.ID})
require.NoError(t, err)
require.Len(t, configs, 1)
require.Equal(t, chatOrgCfg.ID, configs[0].ID)
})
t.Run("ThirdOrgConfigExcluded", func(t *testing.T) {
t.Parallel()
db, _ := dbtestutil.NewDB(t)
ctx := testutil.Context(t, testutil.WaitShort)
chatOrg, chatOrgCfg := newOrgWithConfig(t, db)
_, foreignCfg := newOrgWithConfig(t, db)
configs, err := enabledMCPServerConfigsForChatOrg(ctx, db, chatOrg.ID, []uuid.UUID{chatOrgCfg.ID, foreignCfg.ID})
require.NoError(t, err)
require.Len(t, configs, 1)
require.Equal(t, chatOrgCfg.ID, configs[0].ID)
})
t.Run("DisabledConfigExcluded", func(t *testing.T) {
t.Parallel()
db, _ := dbtestutil.NewDB(t)
ctx := testutil.Context(t, testutil.WaitShort)
// dbgen.MCPServerConfig defaults Enabled to true, so insert the
// disabled config directly.
chatOrg := dbgen.Organization(t, db, database.Organization{})
user := dbgen.User(t, db, database.User{})
disabledCfg, err := db.InsertMCPServerConfig(ctx, database.InsertMCPServerConfigParams{
ID: uuid.New(),
OrganizationID: chatOrg.ID,
DisplayName: "Disabled MCP Server",
Slug: testutil.GetRandomName(t),
Url: "https://mcp.example.com",
Transport: "streamable_http",
AuthType: "none",
ToolAllowList: []string{},
ToolDenyList: []string{},
Availability: "default_off",
Enabled: false,
CreatedBy: user.ID,
UpdatedBy: user.ID,
})
require.NoError(t, err)
configs, err := enabledMCPServerConfigsForChatOrg(ctx, db, chatOrg.ID, []uuid.UUID{disabledCfg.ID})
require.NoError(t, err)
require.Empty(t, configs)
})
t.Run("DuplicateIDsYieldOneConfigPerUniqueID", func(t *testing.T) {
t.Parallel()
db, _ := dbtestutil.NewDB(t)
ctx := testutil.Context(t, testutil.WaitShort)
// The ID array may contain duplicates, but the query returns one row
// per unique ID ordered by display_name.
chatOrg, cfgA := newOrgWithConfig(t, db)
cfgB := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{
OrganizationID: chatOrg.ID,
Enabled: true,
})
// The requested order is the reverse of display_name order to
// prove the output ordering comes from the SQL, not the request.
requested := []uuid.UUID{cfgB.ID, cfgA.ID, cfgB.ID, cfgA.ID}
configs, err := enabledMCPServerConfigsForChatOrg(ctx, db, chatOrg.ID, requested)
require.NoError(t, err)
require.Len(t, configs, 2)
gotIDs := []uuid.UUID{configs[0].ID, configs[1].ID}
wantOrder := []uuid.UUID{cfgA.ID, cfgB.ID}
if cfgA.DisplayName > cfgB.DisplayName {
wantOrder = []uuid.UUID{cfgB.ID, cfgA.ID}
}
require.Equal(t, wantOrder, gotIDs, "output must follow display_name order, not request order")
})
t.Run("ChatOrgWithNoConfigs", func(t *testing.T) {
t.Parallel()
db, _ := dbtestutil.NewDB(t)
ctx := testutil.Context(t, testutil.WaitShort)
defaultOrg, err := db.GetDefaultOrganization(ctx)
require.NoError(t, err)
chatOrg := dbgen.Organization(t, db, database.Organization{})
defaultOrgCfg := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{
OrganizationID: defaultOrg.ID,
Enabled: true,
})
configs, err := enabledMCPServerConfigsForChatOrg(ctx, db, chatOrg.ID, []uuid.UUID{defaultOrgCfg.ID})
require.NoError(t, err)
require.Empty(t, configs)
})
}
+1 -1
View File
@@ -1205,7 +1205,7 @@ func (p *Server) resolveExploreToolSnapshot(
) ([]uuid.UUID, error) {
inheritedMCPServerIDs := []uuid.UUID{}
if len(parent.MCPServerIDs) > 0 {
configs, err := p.db.GetMCPServerConfigsByIDs(ctx, parent.MCPServerIDs)
configs, err := enabledMCPServerConfigsForChatOrg(ctx, p.db, parent.OrganizationID, parent.MCPServerIDs)
if err != nil {
return nil, xerrors.Errorf("get parent MCP server configs for chat %s: %w", parent.ID, err)
}
+38 -31
View File
@@ -693,6 +693,7 @@ func insertInternalChatModelConfigWithOptions(
func insertInternalMCPServerConfig(
t *testing.T,
db database.Store,
organizationID uuid.UUID,
userID uuid.UUID,
slug string,
allowInPlanMode bool,
@@ -700,6 +701,7 @@ func insertInternalMCPServerConfig(
t.Helper()
return dbgen.MCPServerConfig(t, db, database.MCPServerConfig{
OrganizationID: organizationID,
DisplayName: slug,
Slug: slug,
Url: "https://" + slug + ".example.com",
@@ -2379,28 +2381,30 @@ func TestResolveExploreToolSnapshot(t *testing.T) {
db, ps := dbtestutil.NewDB(t)
server := newInternalTestServer(t, db, ps, chatprovider.ProviderAPIKeys{})
user, _, _ := seedInternalChatDeps(t, db)
user, org, _ := seedInternalChatDeps(t, db)
approvedMCP := insertInternalMCPServerConfig(
t, db, user.ID, "approved-"+uuid.NewString(), true,
t, db, org.ID, user.ID, "approved-"+uuid.NewString(), true,
)
blockedMCP := insertInternalMCPServerConfig(
t, db, user.ID, "blocked-"+uuid.NewString(), false,
t, db, org.ID, user.ID, "blocked-"+uuid.NewString(), false,
)
// Build parent chats in memory rather than via server.CreateChat.
// resolveExploreToolSnapshot only reads ID, MCPServerIDs, PlanMode,
// ParentChatID, and Mode from its parent argument, so persisting
// the chats is unnecessary. Skipping CreateChat avoids waking the
// background acquireLoop, which would otherwise try to dial the
// fake MCP URLs and call OpenAI with the dbgen test API key. Those
// side effects were the root cause of the flake tracked in
// CODAGT-367.
// resolveExploreToolSnapshot only reads ID, OrganizationID,
// MCPServerIDs, PlanMode, ParentChatID, and Mode from its parent
// argument, so persisting the chats is unnecessary. Skipping
// CreateChat avoids waking the background acquireLoop, which would
// otherwise try to dial the fake MCP URLs and call OpenAI with the
// dbgen test API key. Those side effects were the root cause of the
// flake tracked in CODAGT-367.
askParent := database.Chat{
ID: uuid.New(),
MCPServerIDs: []uuid.UUID{approvedMCP.ID, blockedMCP.ID},
ID: uuid.New(),
OrganizationID: org.ID,
MCPServerIDs: []uuid.UUID{approvedMCP.ID, blockedMCP.ID},
}
planParent := database.Chat{
ID: uuid.New(),
ID: uuid.New(),
OrganizationID: org.ID,
PlanMode: database.NullChatPlanMode{
ChatPlanMode: database.ChatPlanModePlan,
Valid: true,
@@ -2473,7 +2477,7 @@ func TestCreateChildSubagentChatWithOptions_ExplorePersistsMCPSnapshot(t *testin
ctx, t, server, db, org.ID, user.ID, model.ID, "parent-explore-snapshot",
)
mcpCfg := insertInternalMCPServerConfig(
t, db, user.ID, "snapshot-"+uuid.NewString(), false,
t, db, org.ID, user.ID, "snapshot-"+uuid.NewString(), false,
)
child, err := server.createChildSubagentChatWithOptions(
@@ -2505,10 +2509,10 @@ func TestSpawnAgent_ExploreSnapshotsTurnStateParentState(t *testing.T) {
ctx := chatdTestContext(t)
user, org, model := seedInternalChatDeps(t, db)
turnStartConfig := insertInternalMCPServerConfig(
t, db, user.ID, "turn-start-"+uuid.NewString(), false,
t, db, org.ID, user.ID, "turn-start-"+uuid.NewString(), false,
)
mutatedConfig := insertInternalMCPServerConfig(
t, db, user.ID, "mutated-"+uuid.NewString(), true,
t, db, org.ID, user.ID, "mutated-"+uuid.NewString(), true,
)
parent, err := server.CreateChat(ctx, CreateOptions{
@@ -3359,11 +3363,12 @@ func TestSpawnAgent_ComputerUseInheritsMCPServerIDs(t *testing.T) {
insertEnabledAnthropicProvider(t, db, user.ID)
mcpCfg := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{
DisplayName: "MCP Test",
Slug: "mcp-test",
Url: "https://mcp.example.com",
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
OrganizationID: org.ID,
DisplayName: "MCP Test",
Slug: "mcp-test",
Url: "https://mcp.example.com",
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
})
parentMCPIDs := []uuid.UUID{mcpCfg.ID}
@@ -3410,19 +3415,21 @@ func TestCreateChildSubagentChat_InheritsMCPServerIDs(t *testing.T) {
// Insert two MCP server configs so we can verify both are
// inherited by the child chat.
mcpA := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{
DisplayName: "MCP A",
Slug: "mcp-a",
Url: "https://mcp-a.example.com",
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
OrganizationID: org.ID,
DisplayName: "MCP A",
Slug: "mcp-a",
Url: "https://mcp-a.example.com",
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
})
mcpB := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{
DisplayName: "MCP B",
Slug: "mcp-b",
Url: "https://mcp-b.example.com",
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
OrganizationID: org.ID,
DisplayName: "MCP B",
Slug: "mcp-b",
Url: "https://mcp-b.example.com",
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
})
parentMCPIDs := []uuid.UUID{mcpA.ID, mcpB.ID}