Thanks to visit codestin.com
Credit goes to github.com

Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
72 changes: 49 additions & 23 deletions coderd/database/dbauthz/dbauthz.go
Original file line number Diff line number Diff line change
Expand Up @@ -494,6 +494,29 @@ var (
}.WithCachedASTValue()
}

// MCP token refresh writes use this owner-scoped subject so the shared
// chatd context cannot update arbitrary users' personal data. Reads use
// the daemon-wide chatd subject.
subjectChatdTokenOwner = func(userID uuid.UUID) rbac.Subject {
return rbac.Subject{
Type: rbac.SubjectTypeChatdTokenOwner,
FriendlyName: "Chatd Token Owner",
ID: userID.String(),
Roles: rbac.Roles([]rbac.Role{
{
Identifier: rbac.RoleIdentifier{Name: "chatdtokenowner"},
DisplayName: "Chatd Token Owner",
Site: []rbac.Permission{},
User: rbac.Permissions(map[string][]policy.Action{
rbac.ResourceUser.Type: {policy.ActionUpdatePersonal},
}),
ByOrgID: map[string]rbac.OrgPermissions{},
},
}),
Scope: rbac.ScopeAll,
}.WithCachedASTValue()
}

subjectSystemRestricted = rbac.Subject{
Type: rbac.SubjectTypeSystemRestricted,
FriendlyName: "System",
Expand Down Expand Up @@ -795,7 +818,10 @@ var (
rbac.ResourceWorkspace.Type: {policy.ActionRead, policy.ActionUpdate},
rbac.ResourceDeploymentConfig.Type: {policy.ActionRead},
rbac.ResourceMCPServerConfig.Type: {policy.ActionRead},
rbac.ResourceUser.Type: {policy.ActionReadPersonal},
// Site-wide UpdatePersonal would let chatd write any
// user's personal data; token writes use the per-user
// AsChatdTokenOwner subject instead.
rbac.ResourceUser.Type: {policy.ActionReadPersonal},
}),
User: []rbac.Permission{},
ByOrgID: map[string]rbac.OrgPermissions{},
Expand Down Expand Up @@ -930,6 +956,12 @@ func AsChatdKeyMinter(ctx context.Context, userID uuid.UUID) context.Context {
return As(ctx, subjectChatdKeyMinter(userID))
}

// AsChatdTokenOwner returns a context with an actor that persists MCP
// OAuth2 token refresh results for the specified token owner only.
func AsChatdTokenOwner(ctx context.Context, userID uuid.UUID) context.Context {
return As(ctx, subjectChatdTokenOwner(userID))
}

// AsSystemRestricted returns a context with an actor that has permissions
// required for various system operations (login, logout, metrics cache).
// DO NOT USE THIS UNLESS YOU HAVE ABSOLUTELY NO OTHER CHOICE. Prefer using a
Expand Down Expand Up @@ -2295,10 +2327,9 @@ func (q *querier) DeleteMCPServerConfigByID(ctx context.Context, id uuid.UUID) e
}

func (q *querier) DeleteMCPServerUserToken(ctx context.Context, arg database.DeleteMCPServerUserTokenParams) error {
if err := q.authorizeContext(ctx, policy.ActionUpdate, rbac.ResourceDeploymentConfig); err != nil {
return err
}
return q.db.DeleteMCPServerUserToken(ctx, arg)
return fetchAndExec(q.log, q.auth, policy.ActionUpdatePersonal, func(ctx context.Context, arg database.DeleteMCPServerUserTokenParams) (database.MCPServerUserToken, error) {
return q.db.GetMCPServerUserToken(ctx, database.GetMCPServerUserTokenParams(arg))
}, q.db.DeleteMCPServerUserToken)(ctx, arg)
}

func (q *querier) DeleteMCPServerUserTokensByConfigID(ctx context.Context, mcpServerConfigID uuid.UUID) error {
Expand Down Expand Up @@ -4082,17 +4113,15 @@ func (q *querier) GetMCPServerConfigsByOrganization(ctx context.Context, organiz
}

func (q *querier) GetMCPServerUserToken(ctx context.Context, arg database.GetMCPServerUserTokenParams) (database.MCPServerUserToken, error) {
if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceDeploymentConfig); err != nil {
return database.MCPServerUserToken{}, err
}
return q.db.GetMCPServerUserToken(ctx, arg)
return fetchWithAction(q.log, q.auth, policy.ActionReadPersonal, q.db.GetMCPServerUserToken)(ctx, arg)
}

func (q *querier) GetMCPServerUserTokenByID(ctx context.Context, id uuid.UUID) (database.MCPServerUserToken, error) {
return fetchWithAction(q.log, q.auth, policy.ActionReadPersonal, q.db.GetMCPServerUserTokenByID)(ctx, id)
}

func (q *querier) GetMCPServerUserTokensByUserID(ctx context.Context, userID uuid.UUID) ([]database.MCPServerUserToken, error) {
if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceDeploymentConfig); err != nil {
return nil, err
}
return q.db.GetMCPServerUserTokensByUserID(ctx, userID)
return fetchWithPostFilter(q.auth, policy.ActionReadPersonal, q.db.GetMCPServerUserTokensByUserID)(ctx, userID)
}

func (q *querier) GetNextPendingWorkspaceBuildOrchestrationForUpdate(ctx context.Context) (database.WorkspaceBuildOrchestration, error) {
Expand Down Expand Up @@ -6990,10 +7019,10 @@ func (q *querier) MarkChatsContextDirtyByAgent(ctx context.Context, arg database
}

func (q *querier) MarkMCPServerUserTokenRefreshFailure(ctx context.Context, arg database.MarkMCPServerUserTokenRefreshFailureParams) (database.MCPServerUserToken, error) {
if err := q.authorizeContext(ctx, policy.ActionUpdate, rbac.ResourceDeploymentConfig); err != nil {
return database.MCPServerUserToken{}, err
fetch := func(ctx context.Context, arg database.MarkMCPServerUserTokenRefreshFailureParams) (database.MCPServerUserToken, error) {
return q.db.GetMCPServerUserTokenByID(ctx, arg.ID)
}
return q.db.MarkMCPServerUserTokenRefreshFailure(ctx, arg)
return fetchAndQuery(q.log, q.auth, policy.ActionUpdatePersonal, fetch, q.db.MarkMCPServerUserTokenRefreshFailure)(ctx, arg)
}

func (q *querier) OIDCClaimFieldValues(ctx context.Context, args database.OIDCClaimFieldValuesParams) ([]string, error) {
Expand Down Expand Up @@ -7685,10 +7714,10 @@ func (q *querier) UpdateMCPServerConfigACLByID(ctx context.Context, arg database
}

func (q *querier) UpdateMCPServerUserTokenFromRefresh(ctx context.Context, arg database.UpdateMCPServerUserTokenFromRefreshParams) (database.MCPServerUserToken, error) {
if err := q.authorizeContext(ctx, policy.ActionUpdate, rbac.ResourceDeploymentConfig); err != nil {
return database.MCPServerUserToken{}, err
fetch := func(ctx context.Context, arg database.UpdateMCPServerUserTokenFromRefreshParams) (database.MCPServerUserToken, error) {
return q.db.GetMCPServerUserTokenByID(ctx, arg.ID)
}
return q.db.UpdateMCPServerUserTokenFromRefresh(ctx, arg)
return fetchAndQuery(q.log, q.auth, policy.ActionUpdatePersonal, fetch, q.db.UpdateMCPServerUserTokenFromRefresh)(ctx, arg)
}

func (q *querier) UpdateMemberRoles(ctx context.Context, arg database.UpdateMemberRolesParams) (database.OrganizationMember, error) {
Expand Down Expand Up @@ -9040,10 +9069,7 @@ func (q *querier) UpsertLogoURL(ctx context.Context, value string) error {
}

func (q *querier) UpsertMCPServerUserToken(ctx context.Context, arg database.UpsertMCPServerUserTokenParams) (database.MCPServerUserToken, error) {
if err := q.authorizeContext(ctx, policy.ActionUpdate, rbac.ResourceDeploymentConfig); err != nil {
return database.MCPServerUserToken{}, err
}
return q.db.UpsertMCPServerUserToken(ctx, arg)
return insertWithAction(q.log, q.auth, rbac.ResourceUserObject(arg.UserID), policy.ActionUpdatePersonal, q.db.UpsertMCPServerUserToken)(ctx, arg)
}

func (q *querier) UpsertNotificationReportGeneratorLog(ctx context.Context, arg database.UpsertNotificationReportGeneratorLogParams) error {
Expand Down
43 changes: 31 additions & 12 deletions coderd/database/dbauthz/dbauthz_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -1638,13 +1638,15 @@ func (s *MethodTestSuite) TestChats() {
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) {
s.Run("DeleteMCPServerUserToken", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
arg := database.DeleteMCPServerUserTokenParams{
MCPServerConfigID: uuid.New(),
UserID: uuid.New(),
}
token := testutil.Fake(s.T(), faker, database.MCPServerUserToken{MCPServerConfigID: arg.MCPServerConfigID, UserID: arg.UserID})
dbm.EXPECT().GetMCPServerUserToken(gomock.Any(), database.GetMCPServerUserTokenParams(arg)).Return(token, nil).AnyTimes()
dbm.EXPECT().DeleteMCPServerUserToken(gomock.Any(), arg).Return(nil).AnyTimes()
check.Args(arg).Asserts(rbac.ResourceDeploymentConfig, policy.ActionUpdate)
check.Args(arg).Asserts(token, policy.ActionUpdatePersonal)
}))
s.Run("DeleteMCPServerUserTokensByConfigID", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
config := testutil.Fake(s.T(), faker, database.MCPServerConfig{})
Expand Down Expand Up @@ -1716,13 +1718,19 @@ func (s *MethodTestSuite) TestChats() {
}
token := testutil.Fake(s.T(), faker, database.MCPServerUserToken{MCPServerConfigID: arg.MCPServerConfigID, UserID: arg.UserID})
dbm.EXPECT().GetMCPServerUserToken(gomock.Any(), arg).Return(token, nil).AnyTimes()
check.Args(arg).Asserts(rbac.ResourceDeploymentConfig, policy.ActionRead).Returns(token)
check.Args(arg).Asserts(token, policy.ActionReadPersonal).Returns(token)
}))
s.Run("GetMCPServerUserTokenByID", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
token := testutil.Fake(s.T(), faker, database.MCPServerUserToken{})
dbm.EXPECT().GetMCPServerUserTokenByID(gomock.Any(), token.ID).Return(token, nil).AnyTimes()
check.Args(token.ID).Asserts(token, policy.ActionReadPersonal).Returns(token)
}))
s.Run("GetMCPServerUserTokensByUserID", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
userID := uuid.New()
tokens := []database.MCPServerUserToken{testutil.Fake(s.T(), faker, database.MCPServerUserToken{UserID: userID})}
dbm.EXPECT().GetMCPServerUserTokensByUserID(gomock.Any(), userID).Return(tokens, nil).AnyTimes()
check.Args(userID).Asserts(rbac.ResourceDeploymentConfig, policy.ActionRead).Returns(tokens)
tokenA := testutil.Fake(s.T(), faker, database.MCPServerUserToken{UserID: userID})
tokenB := testutil.Fake(s.T(), faker, database.MCPServerUserToken{UserID: userID})
dbm.EXPECT().GetMCPServerUserTokensByUserID(gomock.Any(), userID).Return([]database.MCPServerUserToken{tokenA, tokenB}, nil).AnyTimes()
check.Args(userID).Asserts(tokenA, policy.ActionReadPersonal, tokenB, policy.ActionReadPersonal).OutOfOrder().Returns([]database.MCPServerUserToken{tokenA, tokenB})
}))
s.Run("InsertMCPServerConfig", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
arg := database.InsertMCPServerConfigParams{
Expand Down Expand Up @@ -1803,8 +1811,9 @@ func (s *MethodTestSuite) TestChats() {
AccessToken: "refreshed-access-token",
TokenType: "bearer",
}
dbm.EXPECT().GetMCPServerUserTokenByID(gomock.Any(), token.ID).Return(token, nil).AnyTimes()
dbm.EXPECT().UpdateMCPServerUserTokenFromRefresh(gomock.Any(), arg).Return(token, nil).AnyTimes()
check.Args(arg).Asserts(rbac.ResourceDeploymentConfig, policy.ActionUpdate).Returns(token)
check.Args(arg).Asserts(token, policy.ActionUpdatePersonal).Returns(token)
}))
s.Run("UpsertMCPServerUserToken", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
arg := database.UpsertMCPServerUserTokenParams{
Expand All @@ -1815,7 +1824,7 @@ func (s *MethodTestSuite) TestChats() {
}
token := testutil.Fake(s.T(), faker, database.MCPServerUserToken{MCPServerConfigID: arg.MCPServerConfigID, UserID: arg.UserID})
dbm.EXPECT().UpsertMCPServerUserToken(gomock.Any(), arg).Return(token, nil).AnyTimes()
check.Args(arg).Asserts(rbac.ResourceDeploymentConfig, policy.ActionUpdate).Returns(token)
check.Args(arg).Asserts(rbac.ResourceUserObject(arg.UserID), policy.ActionUpdatePersonal).Returns(token)
}))
s.Run("MarkMCPServerUserTokenRefreshFailure", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
token := testutil.Fake(s.T(), faker, database.MCPServerUserToken{})
Expand All @@ -1824,8 +1833,9 @@ func (s *MethodTestSuite) TestChats() {
UpdatedAt: token.UpdatedAt,
OauthRefreshFailureReason: "invalid_grant",
}
dbm.EXPECT().GetMCPServerUserTokenByID(gomock.Any(), token.ID).Return(token, nil).AnyTimes()
dbm.EXPECT().MarkMCPServerUserTokenRefreshFailure(gomock.Any(), arg).Return(token, nil).AnyTimes()
check.Args(arg).Asserts(rbac.ResourceDeploymentConfig, policy.ActionUpdate).Returns(token)
check.Args(arg).Asserts(token, policy.ActionUpdatePersonal).Returns(token)
}))
}

Expand Down Expand Up @@ -7602,9 +7612,18 @@ func TestAsChatd(t *testing.T) {
err = auth.Authorize(ctx, actor, policy.ActionUpdate, rbac.ResourceDeploymentConfig)
require.Error(t, err, "deployment config update should not be allowed")

// User read_personal (needed for GetUserChatCustomPrompt).
err = auth.Authorize(ctx, actor, policy.ActionReadPersonal, rbac.ResourceUser)
require.NoError(t, err, "user read_personal should be allowed")
// Pin the complete ResourceUser action set: read_personal (user
// chat custom prompts) only. Token refresh persistence uses the
// per-user AsChatdTokenOwner subject, so a future site-wide
// personal-write grant fails here.
for _, action := range rbac.ResourceUser.AvailableActions() {
err := auth.Authorize(ctx, actor, action, rbac.ResourceUser)
if action == policy.ActionReadPersonal {
require.NoError(t, err, "user %s should be allowed", action)
} else {
require.Error(t, err, "user %s should be denied", action)
}
}
})

t.Run("DeniedActions", func(t *testing.T) {
Expand Down
Loading
Loading