diff --git a/coderd/chatd/chatd_test.go b/coderd/chatd/chatd_test.go index 081fcdaf60..31d5c60534 100644 --- a/coderd/chatd/chatd_test.go +++ b/coderd/chatd/chatd_test.go @@ -638,7 +638,7 @@ func TestCreateChatRejectsWhenUsageLimitReached(t *testing.T) { }) require.NoError(t, err) - beforeChats, err := db.GetChatsByOwnerID(ctx, database.GetChatsByOwnerIDParams{ + beforeChats, err := db.GetChats(ctx, database.GetChatsParams{ OwnerID: user.ID, AfterID: uuid.Nil, OffsetOpt: 0, @@ -660,7 +660,7 @@ func TestCreateChatRejectsWhenUsageLimitReached(t *testing.T) { require.Equal(t, int64(100), limitErr.LimitMicros) require.Equal(t, int64(100), limitErr.ConsumedMicros) - afterChats, err := db.GetChatsByOwnerID(ctx, database.GetChatsByOwnerIDParams{ + afterChats, err := db.GetChats(ctx, database.GetChatsParams{ OwnerID: user.ID, AfterID: uuid.Nil, OffsetOpt: 0, @@ -2678,7 +2678,7 @@ func TestComputerUseSubagentToolsAndModel(t *testing.T) { // 6. Verify the child chat has Mode = computer_use in // the DB. - allChats, err := db.GetChatsByOwnerID(ctx, database.GetChatsByOwnerIDParams{ + allChats, err := db.GetChats(ctx, database.GetChatsParams{ OwnerID: user.ID, }) require.NoError(t, err) diff --git a/coderd/chats.go b/coderd/chats.go index 9fdaf7351a..d4f8b618c7 100644 --- a/coderd/chats.go +++ b/coderd/chats.go @@ -189,7 +189,7 @@ func (api *API) listChats(rw http.ResponseWriter, r *http.Request) { return } - params := database.GetChatsByOwnerIDParams{ + params := database.GetChatsParams{ OwnerID: apiKey.UserID, Archived: searchParams.Archived, AfterID: paginationParams.AfterID, @@ -199,7 +199,7 @@ func (api *API) listChats(rw http.ResponseWriter, r *http.Request) { LimitOpt: int32(paginationParams.Limit), } - chats, err := api.Database.GetChatsByOwnerID(ctx, params) + chats, err := api.Database.GetChats(ctx, params) if err != nil { httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{ Message: "Failed to list chats.", diff --git a/coderd/database/dbauthz/dbauthz.go b/coderd/database/dbauthz/dbauthz.go index 407a281c11..8eaa94d2df 100644 --- a/coderd/database/dbauthz/dbauthz.go +++ b/coderd/database/dbauthz/dbauthz.go @@ -2653,8 +2653,12 @@ func (q *querier) GetChatUsageLimitUserOverride(ctx context.Context, userID uuid return q.db.GetChatUsageLimitUserOverride(ctx, userID) } -func (q *querier) GetChatsByOwnerID(ctx context.Context, ownerID database.GetChatsByOwnerIDParams) ([]database.Chat, error) { - return fetchWithPostFilter(q.auth, policy.ActionRead, q.db.GetChatsByOwnerID)(ctx, ownerID) +func (q *querier) GetChats(ctx context.Context, arg database.GetChatsParams) ([]database.Chat, error) { + prep, err := prepareSQLFilter(ctx, q.auth, policy.ActionRead, rbac.ResourceChat.Type) + if err != nil { + return nil, xerrors.Errorf("(dev error) prepare sql filter: %w", err) + } + return q.db.GetAuthorizedChats(ctx, arg, prep) } func (q *querier) GetConnectionLogsOffset(ctx context.Context, arg database.GetConnectionLogsOffsetParams) ([]database.GetConnectionLogsOffsetRow, error) { @@ -6878,3 +6882,7 @@ func (q *querier) ListAuthorizedAIBridgeModels(ctx context.Context, arg database // database.Store interface, so dbauthz needs to implement it. return q.ListAIBridgeModels(ctx, arg) } + +func (q *querier) GetAuthorizedChats(ctx context.Context, arg database.GetChatsParams, _ rbac.PreparedAuthorized) ([]database.Chat, error) { + return q.GetChats(ctx, arg) +} diff --git a/coderd/database/dbauthz/dbauthz_test.go b/coderd/database/dbauthz/dbauthz_test.go index d23fa42c25..6e82cd4932 100644 --- a/coderd/database/dbauthz/dbauthz_test.go +++ b/coderd/database/dbauthz/dbauthz_test.go @@ -618,12 +618,17 @@ func (s *MethodTestSuite) TestChats() { dbm.EXPECT().GetChatProviders(gomock.Any()).Return([]database.ChatProvider{providerA, providerB}, nil).AnyTimes() check.Args().Asserts(rbac.ResourceDeploymentConfig, policy.ActionRead).Returns([]database.ChatProvider{providerA, providerB}) })) - s.Run("GetChatsByOwnerID", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { - c1 := testutil.Fake(s.T(), faker, database.Chat{}) - c2 := testutil.Fake(s.T(), faker, database.Chat{}) - params := database.GetChatsByOwnerIDParams{OwnerID: c1.OwnerID} - dbm.EXPECT().GetChatsByOwnerID(gomock.Any(), params).Return([]database.Chat{c1, c2}, nil).AnyTimes() - check.Args(params).Asserts(c1, policy.ActionRead, c2, policy.ActionRead).Returns([]database.Chat{c1, c2}) + s.Run("GetChats", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) { + params := database.GetChatsParams{} + dbm.EXPECT().GetAuthorizedChats(gomock.Any(), params, gomock.Any()).Return([]database.Chat{}, nil).AnyTimes() + // No asserts here because SQLFilter. + check.Args(params).Asserts() + })) + s.Run("GetAuthorizedChats", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) { + params := database.GetChatsParams{} + dbm.EXPECT().GetAuthorizedChats(gomock.Any(), params, gomock.Any()).Return([]database.Chat{}, nil).AnyTimes() + // No asserts here because it re-routes through GetChats which uses SQLFilter. + check.Args(params, emptyPreparedAuthorized{}).Asserts() })) s.Run("GetChatQueuedMessages", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { chat := testutil.Fake(s.T(), faker, database.Chat{}) diff --git a/coderd/database/dbmetrics/querymetrics.go b/coderd/database/dbmetrics/querymetrics.go index a6f9a2f20b..471c549fd1 100644 --- a/coderd/database/dbmetrics/querymetrics.go +++ b/coderd/database/dbmetrics/querymetrics.go @@ -1191,11 +1191,11 @@ func (m queryMetricsStore) GetChatUsageLimitUserOverride(ctx context.Context, us return r0, r1 } -func (m queryMetricsStore) GetChatsByOwnerID(ctx context.Context, ownerID database.GetChatsByOwnerIDParams) ([]database.Chat, error) { +func (m queryMetricsStore) GetChats(ctx context.Context, arg database.GetChatsParams) ([]database.Chat, error) { start := time.Now() - r0, r1 := m.s.GetChatsByOwnerID(ctx, ownerID) - m.queryLatencies.WithLabelValues("GetChatsByOwnerID").Observe(time.Since(start).Seconds()) - m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetChatsByOwnerID").Inc() + r0, r1 := m.s.GetChats(ctx, arg) + m.queryLatencies.WithLabelValues("GetChats").Observe(time.Since(start).Seconds()) + m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetChats").Inc() return r0, r1 } @@ -4893,3 +4893,11 @@ func (m queryMetricsStore) ListAuthorizedAIBridgeModels(ctx context.Context, arg m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "ListAuthorizedAIBridgeModels").Inc() return r0, r1 } + +func (m queryMetricsStore) GetAuthorizedChats(ctx context.Context, arg database.GetChatsParams, prepared rbac.PreparedAuthorized) ([]database.Chat, error) { + start := time.Now() + r0, r1 := m.s.GetAuthorizedChats(ctx, arg, prepared) + m.queryLatencies.WithLabelValues("GetAuthorizedChats").Observe(time.Since(start).Seconds()) + m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetAuthorizedChats").Inc() + return r0, r1 +} diff --git a/coderd/database/dbmock/dbmock.go b/coderd/database/dbmock/dbmock.go index 0fc237a8a2..34f3a8131d 100644 --- a/coderd/database/dbmock/dbmock.go +++ b/coderd/database/dbmock/dbmock.go @@ -1731,6 +1731,21 @@ func (mr *MockStoreMockRecorder) GetAuthorizedAuditLogsOffset(ctx, arg, prepared return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAuthorizedAuditLogsOffset", reflect.TypeOf((*MockStore)(nil).GetAuthorizedAuditLogsOffset), ctx, arg, prepared) } +// GetAuthorizedChats mocks base method. +func (m *MockStore) GetAuthorizedChats(ctx context.Context, arg database.GetChatsParams, prepared rbac.PreparedAuthorized) ([]database.Chat, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetAuthorizedChats", ctx, arg, prepared) + ret0, _ := ret[0].([]database.Chat) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetAuthorizedChats indicates an expected call of GetAuthorizedChats. +func (mr *MockStoreMockRecorder) GetAuthorizedChats(ctx, arg, prepared any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAuthorizedChats", reflect.TypeOf((*MockStore)(nil).GetAuthorizedChats), ctx, arg, prepared) +} + // GetAuthorizedConnectionLogsOffset mocks base method. func (m *MockStore) GetAuthorizedConnectionLogsOffset(ctx context.Context, arg database.GetConnectionLogsOffsetParams, prepared rbac.PreparedAuthorized) ([]database.GetConnectionLogsOffsetRow, error) { m.ctrl.T.Helper() @@ -2166,19 +2181,19 @@ func (mr *MockStoreMockRecorder) GetChatUsageLimitUserOverride(ctx, userID any) return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetChatUsageLimitUserOverride", reflect.TypeOf((*MockStore)(nil).GetChatUsageLimitUserOverride), ctx, userID) } -// GetChatsByOwnerID mocks base method. -func (m *MockStore) GetChatsByOwnerID(ctx context.Context, arg database.GetChatsByOwnerIDParams) ([]database.Chat, error) { +// GetChats mocks base method. +func (m *MockStore) GetChats(ctx context.Context, arg database.GetChatsParams) ([]database.Chat, error) { m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "GetChatsByOwnerID", ctx, arg) + ret := m.ctrl.Call(m, "GetChats", ctx, arg) ret0, _ := ret[0].([]database.Chat) ret1, _ := ret[1].(error) return ret0, ret1 } -// GetChatsByOwnerID indicates an expected call of GetChatsByOwnerID. -func (mr *MockStoreMockRecorder) GetChatsByOwnerID(ctx, arg any) *gomock.Call { +// GetChats indicates an expected call of GetChats. +func (mr *MockStoreMockRecorder) GetChats(ctx, arg any) *gomock.Call { mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetChatsByOwnerID", reflect.TypeOf((*MockStore)(nil).GetChatsByOwnerID), ctx, arg) + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetChats", reflect.TypeOf((*MockStore)(nil).GetChats), ctx, arg) } // GetConnectionLogsOffset mocks base method. diff --git a/coderd/database/gentest/modelqueries_test.go b/coderd/database/gentest/modelqueries_test.go index 1025aaf324..2ecb6d66d3 100644 --- a/coderd/database/gentest/modelqueries_test.go +++ b/coderd/database/gentest/modelqueries_test.go @@ -26,6 +26,7 @@ func TestCustomQueriesSyncedRowScan(t *testing.T) { "GetTemplatesWithFilter": "GetAuthorizedTemplates", "GetWorkspaces": "GetAuthorizedWorkspaces", "GetUsers": "GetAuthorizedUsers", + "GetChats": "GetAuthorizedChats", } // Scan custom diff --git a/coderd/database/modelqueries.go b/coderd/database/modelqueries.go index 5f0de02c56..8bceef79eb 100644 --- a/coderd/database/modelqueries.go +++ b/coderd/database/modelqueries.go @@ -52,6 +52,7 @@ type customQuerier interface { auditLogQuerier connectionLogQuerier aibridgeQuerier + chatQuerier } type templateQuerier interface { @@ -738,6 +739,68 @@ func (q *sqlQuerier) CountAuthorizedConnectionLogs(ctx context.Context, arg Coun return count, nil } +type chatQuerier interface { + GetAuthorizedChats(ctx context.Context, arg GetChatsParams, prepared rbac.PreparedAuthorized) ([]Chat, error) +} + +func (q *sqlQuerier) GetAuthorizedChats(ctx context.Context, arg GetChatsParams, prepared rbac.PreparedAuthorized) ([]Chat, error) { + authorizedFilter, err := prepared.CompileToSQL(ctx, rbac.ConfigChats()) + if err != nil { + return nil, xerrors.Errorf("compile authorized filter: %w", err) + } + + filtered, err := insertAuthorizedFilter(getChats, 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: GetAuthorizedChats :many\n%s", filtered) + rows, err := q.db.QueryContext(ctx, query, + arg.OwnerID, + arg.Archived, + arg.AfterID, + arg.OffsetOpt, + arg.LimitOpt, + ) + if err != nil { + return nil, err + } + defer rows.Close() + var items []Chat + for rows.Next() { + var i Chat + if err := rows.Scan( + &i.ID, + &i.OwnerID, + &i.WorkspaceID, + &i.Title, + &i.Status, + &i.WorkerID, + &i.StartedAt, + &i.HeartbeatAt, + &i.CreatedAt, + &i.UpdatedAt, + &i.ParentChatID, + &i.RootChatID, + &i.LastModelConfigID, + &i.Archived, + &i.LastError, + &i.Mode, + ); 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 +} + type aibridgeQuerier interface { ListAuthorizedAIBridgeInterceptions(ctx context.Context, arg ListAIBridgeInterceptionsParams, prepared rbac.PreparedAuthorized) ([]ListAIBridgeInterceptionsRow, error) CountAuthorizedAIBridgeInterceptions(ctx context.Context, arg CountAIBridgeInterceptionsParams, prepared rbac.PreparedAuthorized) (int64, error) diff --git a/coderd/database/querier.go b/coderd/database/querier.go index 3c271e63af..92f5e18342 100644 --- a/coderd/database/querier.go +++ b/coderd/database/querier.go @@ -252,7 +252,7 @@ type sqlcQuerier interface { GetChatUsageLimitConfig(ctx context.Context) (ChatUsageLimitConfig, error) GetChatUsageLimitGroupOverride(ctx context.Context, groupID uuid.UUID) (GetChatUsageLimitGroupOverrideRow, error) GetChatUsageLimitUserOverride(ctx context.Context, userID uuid.UUID) (GetChatUsageLimitUserOverrideRow, error) - GetChatsByOwnerID(ctx context.Context, arg GetChatsByOwnerIDParams) ([]Chat, error) + GetChats(ctx context.Context, arg GetChatsParams) ([]Chat, error) GetConnectionLogsOffset(ctx context.Context, arg GetConnectionLogsOffsetParams) ([]GetConnectionLogsOffsetRow, error) GetCryptoKeyByFeatureAndSequence(ctx context.Context, arg GetCryptoKeyByFeatureAndSequenceParams) (CryptoKey, error) GetCryptoKeys(ctx context.Context) ([]CryptoKey, error) diff --git a/coderd/database/querier_test.go b/coderd/database/querier_test.go index 988e550c4d..5588aa01f1 100644 --- a/coderd/database/querier_test.go +++ b/coderd/database/querier_test.go @@ -1235,6 +1235,230 @@ func TestGetAuthorizedWorkspacesAndAgentsByOwnerID(t *testing.T) { }) } +func TestGetAuthorizedChats(t *testing.T) { + t.Parallel() + if testing.Short() { + t.SkipNow() + } + + sqlDB := testSQLDB(t) + err := migrations.Up(sqlDB) + require.NoError(t, err) + db := database.New(sqlDB) + authorizer := rbac.NewStrictCachingAuthorizer(prometheus.NewRegistry()) + + // Create users with different roles. + owner := dbgen.User(t, db, database.User{ + RBACRoles: []string{rbac.RoleOwner().String()}, + }) + member := dbgen.User(t, db, database.User{}) + secondMember := dbgen.User(t, db, database.User{}) + + // Create FK dependencies: a chat provider and model config. + ctx := testutil.Context(t, testutil.WaitMedium) + _, err = db.InsertChatProvider(ctx, database.InsertChatProviderParams{ + Provider: "openai", + DisplayName: "OpenAI", + APIKey: "test-key", + Enabled: true, + }) + require.NoError(t, err) + + modelCfg, err := db.InsertChatModelConfig(ctx, database.InsertChatModelConfigParams{ + Provider: "openai", + Model: "test-model", + DisplayName: "Test Model", + CreatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true}, + UpdatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true}, + Enabled: true, + IsDefault: true, + ContextLimit: 128000, + CompressionThreshold: 80, + Options: json.RawMessage(`{}`), + }) + require.NoError(t, err) + + // Create 3 chats owned by owner. + for i := range 3 { + _, err := db.InsertChat(ctx, database.InsertChatParams{ + OwnerID: owner.ID, + LastModelConfigID: modelCfg.ID, + Title: fmt.Sprintf("owner chat %d", i+1), + }) + require.NoError(t, err) + } + + // Create 2 chats owned by member. + for i := range 2 { + _, err := db.InsertChat(ctx, database.InsertChatParams{ + OwnerID: member.ID, + LastModelConfigID: modelCfg.ID, + Title: fmt.Sprintf("member chat %d", i+1), + }) + require.NoError(t, err) + } + + t.Run("sqlQuerier", func(t *testing.T) { + t.Parallel() + ctx := testutil.Context(t, testutil.WaitMedium) + + // Member should only see their own 2 chats. + memberSubject, _, err := httpmw.UserRBACSubject(ctx, db, member.ID, rbac.ExpandableScope(rbac.ScopeAll)) + require.NoError(t, err) + preparedMember, err := authorizer.Prepare(ctx, memberSubject, policy.ActionRead, rbac.ResourceChat.Type) + require.NoError(t, err) + memberRows, err := db.GetAuthorizedChats(ctx, database.GetChatsParams{}, preparedMember) + require.NoError(t, err) + require.Len(t, memberRows, 2) + for _, row := range memberRows { + require.Equal(t, member.ID, row.OwnerID, "member should only see own chats") + } + + // Owner should see at least the 5 pre-created chats (site-wide + // access). Parallel subtests may add more, so use GreaterOrEqual. + ownerSubject, _, err := httpmw.UserRBACSubject(ctx, db, owner.ID, rbac.ExpandableScope(rbac.ScopeAll)) + require.NoError(t, err) + preparedOwner, err := authorizer.Prepare(ctx, ownerSubject, policy.ActionRead, rbac.ResourceChat.Type) + require.NoError(t, err) + ownerRows, err := db.GetAuthorizedChats(ctx, database.GetChatsParams{}, preparedOwner) + require.NoError(t, err) + require.GreaterOrEqual(t, len(ownerRows), 5) + + // secondMember has no chats and should see 0. + secondSubject, _, err := httpmw.UserRBACSubject(ctx, db, secondMember.ID, rbac.ExpandableScope(rbac.ScopeAll)) + require.NoError(t, err) + preparedSecond, err := authorizer.Prepare(ctx, secondSubject, policy.ActionRead, rbac.ResourceChat.Type) + require.NoError(t, err) + secondRows, err := db.GetAuthorizedChats(ctx, database.GetChatsParams{}, preparedSecond) + require.NoError(t, err) + require.Len(t, secondRows, 0) + + // Org admin should NOT see other users' chats — chats are + // not org-scoped resources. + orgs, err := db.GetOrganizations(ctx, database.GetOrganizationsParams{}) + require.NoError(t, err) + require.NotEmpty(t, orgs) + orgAdmin := dbgen.User(t, db, database.User{}) + dbgen.OrganizationMember(t, db, database.OrganizationMember{ + UserID: orgAdmin.ID, + OrganizationID: orgs[0].ID, + Roles: []string{rbac.RoleOrgAdmin()}, + }) + orgAdminSubject, _, err := httpmw.UserRBACSubject(ctx, db, orgAdmin.ID, rbac.ExpandableScope(rbac.ScopeAll)) + require.NoError(t, err) + preparedOrgAdmin, err := authorizer.Prepare(ctx, orgAdminSubject, policy.ActionRead, rbac.ResourceChat.Type) + require.NoError(t, err) + orgAdminRows, err := db.GetAuthorizedChats(ctx, database.GetChatsParams{}, preparedOrgAdmin) + require.NoError(t, err) + require.Len(t, orgAdminRows, 0, "org admin with no chats should see 0 chats") + + // OwnerID filter: member queries their own chats. + memberFilterSelf, err := db.GetAuthorizedChats(ctx, database.GetChatsParams{ + OwnerID: member.ID, + }, preparedMember) + require.NoError(t, err) + require.Len(t, memberFilterSelf, 2) + + // OwnerID filter: member queries owner's chats → sees 0. + memberFilterOwner, err := db.GetAuthorizedChats(ctx, database.GetChatsParams{ + OwnerID: owner.ID, + }, preparedMember) + require.NoError(t, err) + require.Len(t, memberFilterOwner, 0) + }) + + t.Run("dbauthz", func(t *testing.T) { + t.Parallel() + ctx := testutil.Context(t, testutil.WaitMedium) + + authzdb := dbauthz.New(db, authorizer, slogtest.Make(t, &slogtest.Options{}), coderdtest.AccessControlStorePointer()) + + // As member: should see only own 2 chats. + memberSubject, _, err := httpmw.UserRBACSubject(ctx, authzdb, member.ID, rbac.ExpandableScope(rbac.ScopeAll)) + require.NoError(t, err) + memberCtx := dbauthz.As(ctx, memberSubject) + memberRows, err := authzdb.GetChats(memberCtx, database.GetChatsParams{}) + require.NoError(t, err) + require.Len(t, memberRows, 2) + for _, row := range memberRows { + require.Equal(t, member.ID, row.OwnerID, "member should only see own chats") + } + + // As owner: should see at least the 5 pre-created chats. + ownerSubject, _, err := httpmw.UserRBACSubject(ctx, authzdb, owner.ID, rbac.ExpandableScope(rbac.ScopeAll)) + require.NoError(t, err) + ownerCtx := dbauthz.As(ctx, ownerSubject) + ownerRows, err := authzdb.GetChats(ownerCtx, database.GetChatsParams{}) + require.NoError(t, err) + require.GreaterOrEqual(t, len(ownerRows), 5) + + // As secondMember: should see 0 chats. + secondSubject, _, err := httpmw.UserRBACSubject(ctx, authzdb, secondMember.ID, rbac.ExpandableScope(rbac.ScopeAll)) + require.NoError(t, err) + secondCtx := dbauthz.As(ctx, secondSubject) + secondRows, err := authzdb.GetChats(secondCtx, database.GetChatsParams{}) + require.NoError(t, err) + require.Len(t, secondRows, 0) + }) + + t.Run("pagination", func(t *testing.T) { + t.Parallel() + ctx := testutil.Context(t, testutil.WaitMedium) + + // Use a dedicated user for pagination to avoid interference + // with the other parallel subtests. + paginationUser := dbgen.User(t, db, database.User{}) + for i := range 7 { + _, err := db.InsertChat(ctx, database.InsertChatParams{ + OwnerID: paginationUser.ID, + LastModelConfigID: modelCfg.ID, + Title: fmt.Sprintf("pagination chat %d", i+1), + }) + require.NoError(t, err) + } + + pagUserSubject, _, err := httpmw.UserRBACSubject(ctx, db, paginationUser.ID, rbac.ExpandableScope(rbac.ScopeAll)) + require.NoError(t, err) + preparedMember, err := authorizer.Prepare(ctx, pagUserSubject, policy.ActionRead, rbac.ResourceChat.Type) + require.NoError(t, err) + + // Fetch first page with limit=2. + page1, err := db.GetAuthorizedChats(ctx, database.GetChatsParams{ + LimitOpt: 2, + }, preparedMember) + require.NoError(t, err) + require.Len(t, page1, 2) + for _, row := range page1 { + require.Equal(t, paginationUser.ID, row.OwnerID, "paginated results must belong to pagination user") + } + + // Fetch remaining pages and collect all chat IDs. + allIDs := make(map[uuid.UUID]struct{}) + for _, row := range page1 { + allIDs[row.ID] = struct{}{} + } + offset := int32(2) + for { + page, err := db.GetAuthorizedChats(ctx, database.GetChatsParams{ + LimitOpt: 2, + OffsetOpt: offset, + }, preparedMember) + require.NoError(t, err) + for _, row := range page { + require.Equal(t, paginationUser.ID, row.OwnerID, "paginated results must belong to pagination user") + allIDs[row.ID] = struct{}{} + } + if len(page) < 2 { + break + } + offset += int32(len(page)) //nolint:gosec // Test code, pagination values are small. + } + + // All 7 member chats should be accounted for with no leakage. + require.Len(t, allIDs, 7, "pagination should return all member chats exactly once") + }) +} + func TestInsertWorkspaceAgentLogs(t *testing.T) { t.Parallel() if testing.Short() { diff --git a/coderd/database/queries.sql.go b/coderd/database/queries.sql.go index ab1c208f0b..ef125ac269 100644 --- a/coderd/database/queries.sql.go +++ b/coderd/database/queries.sql.go @@ -4213,13 +4213,16 @@ func (q *sqlQuerier) GetChatUsageLimitUserOverride(ctx context.Context, userID u return i, err } -const getChatsByOwnerID = `-- name: GetChatsByOwnerID :many +const getChats = `-- name: GetChats :many SELECT id, owner_id, workspace_id, title, status, worker_id, started_at, heartbeat_at, created_at, updated_at, parent_chat_id, root_chat_id, last_model_config_id, archived, last_error, mode FROM chats WHERE - owner_id = $1::uuid + CASE + WHEN $1 :: uuid != '00000000-0000-0000-0000-000000000000'::uuid THEN chats.owner_id = $1 + ELSE true + END AND CASE WHEN $2 :: boolean IS NULL THEN true ELSE chats.archived = $2 :: boolean @@ -4243,6 +4246,8 @@ WHERE ) ELSE true END + -- Authorize Filter clause will be injected below in GetAuthorizedChats + -- @authorize_filter ORDER BY -- Deterministic and consistent ordering of all rows, even if they share -- a timestamp. This is to ensure consistent pagination. @@ -4253,7 +4258,7 @@ LIMIT COALESCE(NULLIF($5 :: int, 0), 50) ` -type GetChatsByOwnerIDParams struct { +type GetChatsParams struct { OwnerID uuid.UUID `db:"owner_id" json:"owner_id"` Archived sql.NullBool `db:"archived" json:"archived"` AfterID uuid.UUID `db:"after_id" json:"after_id"` @@ -4261,8 +4266,8 @@ type GetChatsByOwnerIDParams struct { LimitOpt int32 `db:"limit_opt" json:"limit_opt"` } -func (q *sqlQuerier) GetChatsByOwnerID(ctx context.Context, arg GetChatsByOwnerIDParams) ([]Chat, error) { - rows, err := q.db.QueryContext(ctx, getChatsByOwnerID, +func (q *sqlQuerier) GetChats(ctx context.Context, arg GetChatsParams) ([]Chat, error) { + rows, err := q.db.QueryContext(ctx, getChats, arg.OwnerID, arg.Archived, arg.AfterID, diff --git a/coderd/database/queries/chats.sql b/coderd/database/queries/chats.sql index 9a3e1c8c35..b6f69ed027 100644 --- a/coderd/database/queries/chats.sql +++ b/coderd/database/queries/chats.sql @@ -113,13 +113,16 @@ ORDER BY created_at ASC, id ASC; --- name: GetChatsByOwnerID :many +-- name: GetChats :many SELECT * FROM chats WHERE - owner_id = @owner_id::uuid + CASE + WHEN @owner_id :: uuid != '00000000-0000-0000-0000-000000000000'::uuid THEN chats.owner_id = @owner_id + ELSE true + END AND CASE WHEN sqlc.narg('archived') :: boolean IS NULL THEN true ELSE chats.archived = sqlc.narg('archived') :: boolean @@ -143,6 +146,8 @@ WHERE ) ELSE true END + -- Authorize Filter clause will be injected below in GetAuthorizedChats + -- @authorize_filter ORDER BY -- Deterministic and consistent ordering of all rows, even if they share -- a timestamp. This is to ensure consistent pagination. diff --git a/coderd/gitsync/worker.go b/coderd/gitsync/worker.go index 63a4a00319..ea805da679 100644 --- a/coderd/gitsync/worker.go +++ b/coderd/gitsync/worker.go @@ -48,8 +48,8 @@ type Store interface { UpsertChatDiffStatusReference( ctx context.Context, arg database.UpsertChatDiffStatusReferenceParams, ) (database.ChatDiffStatus, error) - GetChatsByOwnerID( - ctx context.Context, arg database.GetChatsByOwnerIDParams, + GetChats( + ctx context.Context, arg database.GetChatsParams, ) ([]database.Chat, error) } @@ -250,7 +250,7 @@ func (w *Worker) MarkStale( return } - chats, err := w.store.GetChatsByOwnerID(ctx, database.GetChatsByOwnerIDParams{ + chats, err := w.store.GetChats(ctx, database.GetChatsParams{ OwnerID: ownerID, }) if err != nil { diff --git a/coderd/gitsync/worker_test.go b/coderd/gitsync/worker_test.go index 7d27f7a992..07f4e889bb 100644 --- a/coderd/gitsync/worker_test.go +++ b/coderd/gitsync/worker_test.go @@ -469,8 +469,8 @@ func TestWorker_MarkStale_UpsertAndPublish(t *testing.T) { ctrl := gomock.NewController(t) store := dbmock.NewMockStore(ctrl) - store.EXPECT().GetChatsByOwnerID(gomock.Any(), gomock.Any()). - DoAndReturn(func(_ context.Context, arg database.GetChatsByOwnerIDParams) ([]database.Chat, error) { + store.EXPECT().GetChats(gomock.Any(), gomock.Any()). + DoAndReturn(func(_ context.Context, arg database.GetChatsParams) ([]database.Chat, error) { require.Equal(t, ownerID, arg.OwnerID) return []database.Chat{ {ID: chat1, OwnerID: ownerID, WorkspaceID: uuid.NullUUID{UUID: workspaceID, Valid: true}}, @@ -478,13 +478,12 @@ func TestWorker_MarkStale_UpsertAndPublish(t *testing.T) { {ID: chatOther, OwnerID: ownerID, WorkspaceID: uuid.NullUUID{UUID: uuid.New(), Valid: true}}, }, nil }) - store.EXPECT().UpsertChatDiffStatusReference(gomock.Any(), gomock.Any()). - DoAndReturn(func(_ context.Context, arg database.UpsertChatDiffStatusReferenceParams) (database.ChatDiffStatus, error) { - mu.Lock() - upsertRefCalls = append(upsertRefCalls, arg) - mu.Unlock() - return database.ChatDiffStatus{ChatID: arg.ChatID}, nil - }).Times(2) + store.EXPECT().UpsertChatDiffStatusReference(gomock.Any(), gomock.Any()).DoAndReturn(func(_ context.Context, arg database.UpsertChatDiffStatusReferenceParams) (database.ChatDiffStatus, error) { + mu.Lock() + upsertRefCalls = append(upsertRefCalls, arg) + mu.Unlock() + return database.ChatDiffStatus{ChatID: arg.ChatID}, nil + }).Times(2) pub := func(_ context.Context, chatID uuid.UUID) error { mu.Lock() @@ -527,7 +526,7 @@ func TestWorker_MarkStale_NoMatchingChats(t *testing.T) { ctrl := gomock.NewController(t) store := dbmock.NewMockStore(ctrl) - store.EXPECT().GetChatsByOwnerID(gomock.Any(), gomock.Any()). + store.EXPECT().GetChats(gomock.Any(), gomock.Any()). Return([]database.Chat{ {ID: uuid.New(), OwnerID: ownerID, WorkspaceID: uuid.NullUUID{UUID: uuid.New(), Valid: true}}, {ID: uuid.New(), OwnerID: ownerID, WorkspaceID: uuid.NullUUID{UUID: uuid.New(), Valid: true}}, @@ -555,7 +554,7 @@ func TestWorker_MarkStale_UpsertFails_ContinuesNext(t *testing.T) { ctrl := gomock.NewController(t) store := dbmock.NewMockStore(ctrl) - store.EXPECT().GetChatsByOwnerID(gomock.Any(), gomock.Any()). + store.EXPECT().GetChats(gomock.Any(), gomock.Any()). Return([]database.Chat{ {ID: chat1, OwnerID: ownerID, WorkspaceID: uuid.NullUUID{UUID: workspaceID, Valid: true}}, {ID: chat2, OwnerID: ownerID, WorkspaceID: uuid.NullUUID{UUID: workspaceID, Valid: true}}, @@ -590,7 +589,7 @@ func TestWorker_MarkStale_GetChatsFails(t *testing.T) { ctrl := gomock.NewController(t) store := dbmock.NewMockStore(ctrl) - store.EXPECT().GetChatsByOwnerID(gomock.Any(), gomock.Any()). + store.EXPECT().GetChats(gomock.Any(), gomock.Any()). Return(nil, fmt.Errorf("db error")) mClock := quartz.NewMock(t) diff --git a/coderd/rbac/authz.go b/coderd/rbac/authz.go index 218283f381..264970928b 100644 --- a/coderd/rbac/authz.go +++ b/coderd/rbac/authz.go @@ -688,6 +688,15 @@ func ConfigWithoutACL() regosql.ConvertConfig { } } +// ConfigChats is the configuration for converting rego to SQL when +// the target table is "chats", which has no organization_id or ACL +// columns. +func ConfigChats() regosql.ConvertConfig { + return regosql.ConvertConfig{ + VariableConverter: regosql.ChatConverter(), + } +} + func ConfigWorkspaces() regosql.ConvertConfig { return regosql.ConvertConfig{ VariableConverter: regosql.WorkspaceConverter(), diff --git a/coderd/rbac/regosql/compile_test.go b/coderd/rbac/regosql/compile_test.go index 7bea7f76fd..9249e890ad 100644 --- a/coderd/rbac/regosql/compile_test.go +++ b/coderd/rbac/regosql/compile_test.go @@ -282,6 +282,22 @@ neq(input.object.owner, ""); p("'10d03e62-7703-4df5-a358-4f76577d4e2f' = id :: text") + " AND " + p("id :: text != ''") + " AND " + p("'' = ''"), ), }, + { + Name: "ChatOwnerMe", + Queries: []string{ + `"me" = input.object.owner; input.object.owner != ""; input.object.org_owner = ""`, + }, + ExpectedSQL: p(p("'me' = owner_id :: text") + " AND " + p("owner_id :: text != ''") + " AND " + p("'' = ''")), + VariableConverter: regosql.ChatConverter(), + }, + { + Name: "ChatOrgScopedNeverMatches", + Queries: []string{ + `input.object.org_owner = "org-id"`, + }, + ExpectedSQL: p("'' = 'org-id'"), + VariableConverter: regosql.ChatConverter(), + }, } for _, tc := range testCases { diff --git a/coderd/rbac/regosql/configs.go b/coderd/rbac/regosql/configs.go index 355a49756d..4f156e8a26 100644 --- a/coderd/rbac/regosql/configs.go +++ b/coderd/rbac/regosql/configs.go @@ -126,6 +126,30 @@ func NoACLConverter() *sqltypes.VariableConverter { return matcher } +// ChatConverter should be used for the chats table, which has no +// organization_id, group_acl, or user_acl columns. +func ChatConverter() *sqltypes.VariableConverter { + matcher := sqltypes.NewVariableConverter().RegisterMatcher( + resourceIDMatcher(), + // The chats table has no organization_id column. Map org_owner + // to a literal empty string so that: + // - User-level ownership checks (org_owner = '') activate correctly. + // - Org-scoped permissions never match (org_owner will never equal + // a real org UUID), which is intentional since chats are not + // org-scoped resources. + // Note: custom org roles that include "chat" permissions will + // silently have no effect because of this mapping. + sqltypes.StringVarMatcher("''", []string{"input", "object", "org_owner"}), + userOwnerMatcher(), + ) + matcher.RegisterMatcher( + sqltypes.AlwaysFalse(groupACLMatcher(matcher)), + sqltypes.AlwaysFalse(userACLMatcher(matcher)), + ) + + return matcher +} + func DefaultVariableConverter() *sqltypes.VariableConverter { matcher := sqltypes.NewVariableConverter().RegisterMatcher( resourceIDMatcher(), diff --git a/coderd/searchquery/search.go b/coderd/searchquery/search.go index 581b004598..7d8f517d08 100644 --- a/coderd/searchquery/search.go +++ b/coderd/searchquery/search.go @@ -471,8 +471,8 @@ func Tasks(ctx context.Context, db database.Store, query string, actorID uuid.UU // // Supported query parameters: // - archived: boolean (default: false, excludes archived chats unless explicitly set) -func Chats(query string) (database.GetChatsByOwnerIDParams, []codersdk.ValidationError) { - filter := database.GetChatsByOwnerIDParams{ +func Chats(query string) (database.GetChatsParams, []codersdk.ValidationError) { + filter := database.GetChatsParams{ // Default to hiding archived chats. Archived: sql.NullBool{Bool: false, Valid: true}, } diff --git a/coderd/searchquery/search_test.go b/coderd/searchquery/search_test.go index 2f6bfb41b0..8e6013ad5a 100644 --- a/coderd/searchquery/search_test.go +++ b/coderd/searchquery/search_test.go @@ -1222,27 +1222,27 @@ func TestSearchChats(t *testing.T) { testCases := []struct { Name string Query string - Expected database.GetChatsByOwnerIDParams + Expected database.GetChatsParams ExpectedErrorContains string }{ { Name: "Empty", Query: "", - Expected: database.GetChatsByOwnerIDParams{ + Expected: database.GetChatsParams{ Archived: sql.NullBool{Bool: false, Valid: true}, }, }, { Name: "ArchivedTrue", Query: "archived:true", - Expected: database.GetChatsByOwnerIDParams{ + Expected: database.GetChatsParams{ Archived: sql.NullBool{Bool: true, Valid: true}, }, }, { Name: "ArchivedFalse", Query: "archived:false", - Expected: database.GetChatsByOwnerIDParams{ + Expected: database.GetChatsParams{ Archived: sql.NullBool{Bool: false, Valid: true}, }, },