mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
fix: limit shared chats to ACL grants (#26123)
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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.
|
||||
|
||||
Generated
+31
-22
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user