diff --git a/coderd/database/modelqueries.go b/coderd/database/modelqueries.go index 972a104201..c8f27d9702 100644 --- a/coderd/database/modelqueries.go +++ b/coderd/database/modelqueries.go @@ -756,6 +756,9 @@ func (q *sqlQuerier) GetAuthorizedChats(ctx context.Context, arg GetChatsParams, if (arg.OwnedOnly || arg.SharedOnly) && arg.ViewerID == uuid.Nil { return nil, xerrors.New("viewer_id required when owned_only or shared_only is true") } + if arg.SharedOnly && arg.SharedWithUserID == uuid.Nil && len(arg.SharedWithGroupIds) == 0 { + return nil, xerrors.New("shared_with_user_id or shared_with_group_ids required when shared_only is true") + } authorizedFilter, err := prepared.CompileToSQL(ctx, rbac.ConfigChats()) if err != nil { @@ -773,6 +776,8 @@ func (q *sqlQuerier) GetAuthorizedChats(ctx context.Context, arg GetChatsParams, arg.OwnedOnly, arg.ViewerID, arg.SharedOnly, + arg.SharedWithUserID, + pq.Array(arg.SharedWithGroupIds), arg.Archived, arg.AfterID, arg.LabelFilter, diff --git a/coderd/database/querier_test.go b/coderd/database/querier_test.go index bc884a0752..32142917f9 100644 --- a/coderd/database/querier_test.go +++ b/coderd/database/querier_test.go @@ -1387,6 +1387,12 @@ func TestGetAuthorizedChats(t *testing.T) { SharedOnly: true, }, preparedMember) require.ErrorContains(t, err, "viewer_id required") + + _, err = db.GetAuthorizedChats(ctx, database.GetChatsParams{ + SharedOnly: true, + ViewerID: member.ID, + }, preparedMember) + require.ErrorContains(t, err, "shared_with_user_id or shared_with_group_ids required") }) t.Run("dbauthz", func(t *testing.T) { @@ -1414,6 +1420,15 @@ func TestGetAuthorizedChats(t *testing.T) { require.NoError(t, err) require.GreaterOrEqual(t, len(ownerRows), 5) + ownerSharedRows, err := authzdb.GetChats(ownerCtx, database.GetChatsParams{ + SharedOnly: true, + ViewerID: owner.ID, + SharedWithUserID: owner.ID, + SharedWithGroupIds: []string{}, + }) + require.NoError(t, err) + require.Empty(t, ownerSharedRows, "shared-only must not include chats visible through owner RBAC") + // As secondMember: should see 0 chats. secondSubject, _, err := httpmw.UserRBACSubject(ctx, authzdb, secondMember.ID, rbac.ExpandableScope(rbac.ScopeAll)) require.NoError(t, err) @@ -1563,8 +1578,9 @@ func TestGetAuthorizedChatsACLSharing(t *testing.T) { require.ElementsMatch(t, []uuid.UUID{ownerChat.ID, recipientChat.ID}, chatIDs(rows)) sharedOnly, err := db.GetAuthorizedChats(ctx, database.GetChatsParams{ - SharedOnly: true, - ViewerID: recipient.ID, + SharedOnly: true, + ViewerID: recipient.ID, + SharedWithUserID: recipient.ID, }, preparedRecipient) require.NoError(t, err) require.ElementsMatch(t, []uuid.UUID{ownerChat.ID}, chatIDs(sharedOnly)) @@ -1584,6 +1600,14 @@ func TestGetAuthorizedChatsACLSharing(t *testing.T) { require.NoError(t, err) require.ElementsMatch(t, []uuid.UUID{ownerChat.ID, recipientChat.ID}, chatIDs(authzRows)) + authzSharedOnly, err := authzdb.GetChats(recipientCtx, database.GetChatsParams{ + SharedOnly: true, + ViewerID: recipient.ID, + SharedWithUserID: recipient.ID, + }) + require.NoError(t, err) + require.ElementsMatch(t, []uuid.UUID{ownerChat.ID}, chatIDs(authzSharedOnly)) + rbac.SetChatACLDisabled(true) disabledRows, err := db.GetAuthorizedChats(ctx, database.GetChatsParams{}, preparedRecipient) require.NoError(t, err) @@ -1673,14 +1697,26 @@ func TestGetAuthorizedChatsACLSharingGroupACL(t *testing.T) { require.ElementsMatch(t, []uuid.UUID{ownerChat.ID, recipientChat.ID}, chatIDs(rows)) sharedOnly, err := db.GetAuthorizedChats(ctx, database.GetChatsParams{ - SharedOnly: true, - ViewerID: recipient.ID, + SharedOnly: true, + ViewerID: recipient.ID, + SharedWithGroupIds: []string{group.ID.String()}, }, preparedRecipient) require.NoError(t, err) require.Len(t, sharedOnly, 1) require.Equal(t, ownerChat.ID, sharedOnly[0].Chat.ID) require.Empty(t, sharedOnly[0].Chat.UserACL) require.Equal(t, sharedGroupACL, sharedOnly[0].Chat.GroupACL) + + authzdb := dbauthz.New(db, authorizer, slogtest.Make(t, &slogtest.Options{}), coderdtest.AccessControlStorePointer()) + recipientCtx := dbauthz.As(ctx, recipientSubject) + authzSharedOnly, err := authzdb.GetChats(recipientCtx, database.GetChatsParams{ + SharedOnly: true, + ViewerID: recipient.ID, + SharedWithGroupIds: []string{group.ID.String()}, + }) + require.NoError(t, err) + require.Len(t, authzSharedOnly, 1) + require.Equal(t, ownerChat.ID, authzSharedOnly[0].Chat.ID) } //nolint:tparallel,paralleltest // It toggles the global chat ACL flag. diff --git a/coderd/database/queries.sql.go b/coderd/database/queries.sql.go index f7901d6ae1..f5d4bc7cf8 100644 --- a/coderd/database/queries.sql.go +++ b/coderd/database/queries.sql.go @@ -8009,7 +8009,7 @@ WITH cursor_chat AS ( updated_at, id FROM chats - WHERE id = $5 + WHERE id = $7 ) SELECT chats_expanded.id, chats_expanded.owner_id, chats_expanded.workspace_id, chats_expanded.title, chats_expanded.status, chats_expanded.worker_id, chats_expanded.started_at, chats_expanded.heartbeat_at, chats_expanded.created_at, chats_expanded.updated_at, chats_expanded.parent_chat_id, chats_expanded.root_chat_id, chats_expanded.last_model_config_id, chats_expanded.archived, chats_expanded.last_error, chats_expanded.mode, chats_expanded.mcp_server_ids, chats_expanded.labels, chats_expanded.build_id, chats_expanded.agent_id, chats_expanded.pin_order, chats_expanded.last_read_message_id, chats_expanded.last_injected_context, chats_expanded.dynamic_tools, chats_expanded.organization_id, chats_expanded.plan_mode, chats_expanded.client_type, chats_expanded.last_turn_summary, chats_expanded.user_acl, chats_expanded.group_acl, chats_expanded.owner_username, chats_expanded.owner_name, @@ -8028,19 +8028,24 @@ WHERE ELSE true END AND CASE - WHEN $3::boolean THEN chats_expanded.owner_id != $2::uuid + WHEN $3::boolean THEN + chats_expanded.owner_id != $2::uuid + AND ( + chats_expanded.user_acl ? ($4::uuid)::text + OR chats_expanded.group_acl ?| $5::text[] + ) ELSE true END AND CASE - WHEN $4 :: boolean IS NULL THEN true - ELSE chats_expanded.archived = $4 :: boolean + WHEN $6 :: boolean IS NULL THEN true + ELSE chats_expanded.archived = $6 :: boolean END AND CASE -- Cursor pagination: the last element on a page acts as the cursor. -- The 4-tuple matches the ORDER BY below. All columns sort DESC -- (pin_order is negated so lower values sort first in DESC order), -- which lets us use a single tuple < comparison. - WHEN $5 :: uuid != '00000000-0000-0000-0000-000000000000'::uuid THEN ( + WHEN $7 :: uuid != '00000000-0000-0000-0000-000000000000'::uuid THEN ( (CASE WHEN chats_expanded.pin_order > 0 THEN 1 ELSE 0 END, -chats_expanded.pin_order, chats_expanded.updated_at, chats_expanded.id) < ( SELECT CASE WHEN cursor_chat.pin_order > 0 THEN 1 ELSE 0 END, @@ -8054,7 +8059,7 @@ WHERE ELSE true END AND CASE - WHEN $6::jsonb IS NOT NULL THEN chats_expanded.labels @> $6::jsonb + WHEN $8::jsonb IS NOT NULL THEN chats_expanded.labels @> $8::jsonb ELSE true END -- Match chats whose linked diff URL (e.g. a pull request URL) @@ -8062,13 +8067,13 @@ WHERE -- a delegated sub-agent's diff status, so we surface the root chat -- when any descendant matches. AND CASE - WHEN $7::text IS NOT NULL THEN EXISTS ( + WHEN $9::text IS NOT NULL THEN EXISTS ( SELECT 1 FROM chat_diff_statuses cds JOIN chats c2 ON c2.id = cds.chat_id WHERE cds.url IS NOT NULL AND cds.url <> '' - AND LOWER(cds.url) = LOWER($7::text) + AND LOWER(cds.url) = LOWER($9::text) AND (c2.id = chats_expanded.id OR c2.root_chat_id = chats_expanded.id) ) ELSE true @@ -8076,11 +8081,11 @@ WHERE -- Filter by title substring (case-insensitive). Applied when the -- caller provides a non-empty title_query. AND CASE - WHEN $8 :: text != '' THEN chats_expanded.title ILIKE '%' || $8 || '%' + WHEN $10 :: text != '' THEN chats_expanded.title ILIKE '%' || $10 || '%' ELSE true END AND CASE - WHEN $9::boolean IS NOT NULL THEN ( + WHEN $11::boolean IS NOT NULL THEN ( EXISTS ( SELECT 1 FROM chat_messages cm WHERE cm.chat_id = chats_expanded.id @@ -8088,7 +8093,7 @@ WHERE AND cm.deleted = false AND cm.id > COALESCE(chats_expanded.last_read_message_id, 0) ) - ) = $9::boolean + ) = $11::boolean ELSE true END -- Filter by pull request status. Unlike the diff_url filter above, @@ -8097,7 +8102,7 @@ WHERE -- parent, so gitsync populates identical PR state on both; traversing -- descendants would be redundant. AND CASE - WHEN COALESCE(array_length($10::text[], 1), 0) > 0 THEN EXISTS ( + WHEN COALESCE(array_length($12::text[], 1), 0) > 0 THEN EXISTS ( SELECT 1 FROM chat_diff_statuses cds WHERE cds.chat_id = chats_expanded.id @@ -8107,40 +8112,40 @@ WHERE WHEN cds.pull_request_state = 'open' THEN 'open' ELSE cds.pull_request_state END - ) = ANY($10::text[]) + ) = ANY($12::text[]) ) ELSE true END -- Filter by PR number (exact match on chat's diff status). AND CASE - WHEN $11::int != 0 THEN EXISTS ( + WHEN $13::int != 0 THEN EXISTS ( SELECT 1 FROM chat_diff_statuses cds WHERE cds.chat_id = chats_expanded.id - AND cds.pr_number = $11 + AND cds.pr_number = $13 ) ELSE true END -- Filter by repository (substring match on remote origin or PR URL). AND CASE - WHEN $12::text != '' THEN EXISTS ( + WHEN $14::text != '' THEN EXISTS ( SELECT 1 FROM chat_diff_statuses cds WHERE cds.chat_id = chats_expanded.id AND ( - cds.git_remote_origin ILIKE '%' || $12 || '%' - OR cds.url ILIKE '%' || $12 || '%' + cds.git_remote_origin ILIKE '%' || $14 || '%' + OR cds.url ILIKE '%' || $14 || '%' ) ) ELSE true END -- Filter by pull request title (case-insensitive substring). AND CASE - WHEN $13::text != '' THEN EXISTS ( + WHEN $15::text != '' THEN EXISTS ( SELECT 1 FROM chat_diff_statuses cds WHERE cds.chat_id = chats_expanded.id - AND cds.pull_request_title ILIKE '%' || $13 || '%' + AND cds.pull_request_title ILIKE '%' || $15 || '%' ) ELSE true END @@ -8160,17 +8165,19 @@ ORDER BY -chats_expanded.pin_order DESC, chats_expanded.updated_at DESC, chats_expanded.id DESC -OFFSET $14 +OFFSET $16 LIMIT -- The chat list is unbounded and expected to grow large. -- Default to 50 to prevent accidental excessively large queries. - COALESCE(NULLIF($15 :: int, 0), 50) + COALESCE(NULLIF($17 :: int, 0), 50) ` type GetChatsParams struct { OwnedOnly bool `db:"owned_only" json:"owned_only"` ViewerID uuid.UUID `db:"viewer_id" json:"viewer_id"` SharedOnly bool `db:"shared_only" json:"shared_only"` + SharedWithUserID uuid.UUID `db:"shared_with_user_id" json:"shared_with_user_id"` + SharedWithGroupIds []string `db:"shared_with_group_ids" json:"shared_with_group_ids"` Archived sql.NullBool `db:"archived" json:"archived"` AfterID uuid.UUID `db:"after_id" json:"after_id"` LabelFilter pqtype.NullRawMessage `db:"label_filter" json:"label_filter"` @@ -8195,6 +8202,8 @@ func (q *sqlQuerier) GetChats(ctx context.Context, arg GetChatsParams) ([]GetCha arg.OwnedOnly, arg.ViewerID, arg.SharedOnly, + arg.SharedWithUserID, + pq.Array(arg.SharedWithGroupIds), arg.Archived, arg.AfterID, arg.LabelFilter, diff --git a/coderd/database/queries/chats.sql b/coderd/database/queries/chats.sql index c8b6502cf5..106e2cc3e6 100644 --- a/coderd/database/queries/chats.sql +++ b/coderd/database/queries/chats.sql @@ -486,7 +486,12 @@ WHERE ELSE true END AND CASE - WHEN @shared_only::boolean THEN chats_expanded.owner_id != @viewer_id::uuid + WHEN @shared_only::boolean THEN + chats_expanded.owner_id != @viewer_id::uuid + AND ( + chats_expanded.user_acl ? (@shared_with_user_id::uuid)::text + OR chats_expanded.group_acl ?| @shared_with_group_ids::text[] + ) ELSE true END AND CASE diff --git a/coderd/exp_chats.go b/coderd/exp_chats.go index d44c326666..b6b32bd7a9 100644 --- a/coderd/exp_chats.go +++ b/coderd/exp_chats.go @@ -390,10 +390,28 @@ func (api *API) listChats(rw http.ResponseWriter, r *http.Request) { } } + var sharedWithGroupIDs []string + if searchParams.SharedOnly { + groups, err := api.Database.GetGroups(ctx, database.GetGroupsParams{HasMemberID: apiKey.UserID}) + if err != nil { + httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{ + Message: "Failed to list chats.", + Detail: err.Error(), + }) + return + } + sharedWithGroupIDs = make([]string, 0, len(groups)) + for _, group := range groups { + sharedWithGroupIDs = append(sharedWithGroupIDs, group.Group.ID.String()) + } + } + params := database.GetChatsParams{ OwnedOnly: searchParams.OwnedOnly, - SharedOnly: searchParams.SharedOnly, ViewerID: apiKey.UserID, + SharedOnly: searchParams.SharedOnly, + SharedWithUserID: apiKey.UserID, + SharedWithGroupIds: sharedWithGroupIDs, Archived: searchParams.Archived, AfterID: paginationParams.AfterID, LabelFilter: labelFilter,