mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
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:
co-authored by
Mathias Fredriksson
parent
7ff1278ab3
commit
443e3b9b80
Generated
+12
@@ -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",
|
||||
|
||||
Generated
+12
@@ -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
@@ -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))
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
Generated
+83
-39
@@ -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.
|
||||
|
||||
Generated
+14
-3
@@ -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
@@ -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);
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
Vendored
+149
@@ -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'
|
||||
);
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
Generated
+17
-1
@@ -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 {
|
||||
|
||||
Generated
+7
-5
@@ -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
|
||||
|
||||
Generated
+194
-97
@@ -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
|
||||
}
|
||||
|
||||
@@ -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 = (
|
||||
|
||||
Generated
+1
-1
@@ -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
@@ -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
@@ -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",
|
||||
|
||||
@@ -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
@@ -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
File diff suppressed because it is too large
Load Diff
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
},
|
||||
}
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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},
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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}
|
||||
|
||||
Reference in New Issue
Block a user