From b9fad662145f883d0e1c8437a51953eba1a07ba0 Mon Sep 17 00:00:00 2001 From: Susana Ferreira Date: Thu, 23 Jul 2026 10:46:58 +0100 Subject: [PATCH] refactor: authorize AI budget reads against the user resource directly (#27443) Replaces the `GetUserByID` read used as an authz check in the AI budget-resolution queries with a targeted `authorizeContext` against the user resource. Same RBAC decision, one fewer db query per resolution step. Follow-up to https://github.com/coder/coder/pull/27364#discussion_r3632577802. > [!NOTE] > Initially generated by Claude Opus 4.7, modified and reviewed by @ssncferreira --- coderd/database/dbauthz/dbauthz.go | 6 +++--- coderd/database/dbauthz/dbauthz_test.go | 3 --- 2 files changed, 3 insertions(+), 6 deletions(-) diff --git a/coderd/database/dbauthz/dbauthz.go b/coderd/database/dbauthz/dbauthz.go index 3cb39f7879..3af34d6b36 100644 --- a/coderd/database/dbauthz/dbauthz.go +++ b/coderd/database/dbauthz/dbauthz.go @@ -3933,7 +3933,7 @@ func (q *querier) GetHealthSettings(ctx context.Context) (string, error) { } func (q *querier) GetHighestGroupAIBudgetByUser(ctx context.Context, userID uuid.UUID) (database.GetHighestGroupAIBudgetByUserRow, error) { - if _, err := q.GetUserByID(ctx, userID); err != nil { // AuthZ check + if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceUserObject(userID)); err != nil { return database.GetHighestGroupAIBudgetByUserRow{}, err } return q.db.GetHighestGroupAIBudgetByUser(ctx, userID) @@ -4940,7 +4940,7 @@ func (q *querier) GetUnexpiredLicenses(ctx context.Context) ([]database.License, } func (q *querier) GetUserAIBudgetOverride(ctx context.Context, userID uuid.UUID) (database.UserAIBudgetOverride, error) { - if _, err := q.GetUserByID(ctx, userID); err != nil { // AuthZ check + if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceUserObject(userID)); err != nil { return database.UserAIBudgetOverride{}, err } return q.db.GetUserAIBudgetOverride(ctx, userID) @@ -5112,7 +5112,7 @@ func (q *querier) GetUserCount(ctx context.Context, includeSystem bool) (int64, } func (q *querier) GetUserEveryoneFallbackGroup(ctx context.Context, userID uuid.UUID) (uuid.UUID, error) { - if _, err := q.GetUserByID(ctx, userID); err != nil { // AuthZ check + if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceUserObject(userID)); err != nil { return uuid.Nil, err } return q.db.GetUserEveryoneFallbackGroup(ctx, userID) diff --git a/coderd/database/dbauthz/dbauthz_test.go b/coderd/database/dbauthz/dbauthz_test.go index d064aeb9bf..8f48b870be 100644 --- a/coderd/database/dbauthz/dbauthz_test.go +++ b/coderd/database/dbauthz/dbauthz_test.go @@ -6998,7 +6998,6 @@ func (s *MethodTestSuite) TestAIBridge() { s.Run("GetUserAIBudgetOverride", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { user := testutil.Fake(s.T(), faker, database.User{}) override := testutil.Fake(s.T(), faker, database.UserAIBudgetOverride{UserID: user.ID}) - dbm.EXPECT().GetUserByID(gomock.Any(), user.ID).Return(user, nil).AnyTimes() dbm.EXPECT().GetUserAIBudgetOverride(gomock.Any(), user.ID).Return(override, nil).AnyTimes() check.Args(user.ID).Asserts(user, policy.ActionRead).Returns(override) })) @@ -7006,7 +7005,6 @@ func (s *MethodTestSuite) TestAIBridge() { s.Run("GetHighestGroupAIBudgetByUser", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { user := testutil.Fake(s.T(), faker, database.User{}) row := testutil.Fake(s.T(), faker, database.GetHighestGroupAIBudgetByUserRow{}) - dbm.EXPECT().GetUserByID(gomock.Any(), user.ID).Return(user, nil).AnyTimes() dbm.EXPECT().GetHighestGroupAIBudgetByUser(gomock.Any(), user.ID).Return(row, nil).AnyTimes() check.Args(user.ID).Asserts(user, policy.ActionRead).Returns(row) })) @@ -7014,7 +7012,6 @@ func (s *MethodTestSuite) TestAIBridge() { s.Run("GetUserEveryoneFallbackGroup", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { user := testutil.Fake(s.T(), faker, database.User{}) group := testutil.Fake(s.T(), faker, database.Group{}) - dbm.EXPECT().GetUserByID(gomock.Any(), user.ID).Return(user, nil).AnyTimes() dbm.EXPECT().GetUserEveryoneFallbackGroup(gomock.Any(), user.ID).Return(group.ID, nil).AnyTimes() check.Args(user.ID).Asserts(user, policy.ActionRead).Returns(group.ID) }))