fix: limit shared chats to ACL grants (#26123)

This commit is contained in:
Danielle Maywood
2026-06-08 12:21:16 +01:00
committed by GitHub
parent 8ae1a5c766
commit 1e5dd83a95
5 changed files with 101 additions and 28 deletions
+5
View File
@@ -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,
+40 -4
View File
@@ -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.
+31 -22
View File
@@ -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,
+6 -1
View File
@@ -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