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) }))