fix: exclude subagent chats from sidebar pagination (#24404)

GetChats now returns only root chats (parent_chat_id IS NULL).
A new GetChildChatsByParentIDs query fetches children for visible
roots and embeds them in each parent's Children field. The
singular getChat endpoint does the same.

Archive invariant is one-way: parent archived implies child
archived. Parent archive/unarchive cascades via root_chat_id.
Individual child archive is permitted; child unarchive while the
parent is archived is rejected atomically (row lock on child,
re-read parent inside the transaction). Embedded children are
filtered by the caller's archive state so individually-archived
children stay hidden from active-parent views.

Gitsync MarkStale uses GetChatsByWorkspaceIDs directly;
MarkStaleParams.OwnerID removed (dead after the switch).

Frontend: buildChatTree reads from the embedded children field,
WebSocket handlers route child events into the parent's children
array, and archiving a child strips it from the parent cache.
This commit is contained in:
Mathias Fredriksson
2026-04-20 13:19:59 +03:00
committed by GitHub
parent df429b7f60
commit fc2493780f
30 changed files with 1514 additions and 225 deletions
+58 -9
View File
@@ -1462,20 +1462,69 @@ func (p *Server) ArchiveChat(ctx context.Context, chat database.Chat) error {
return nil
}
// UnarchiveChat unarchives a chat family and publishes created events for
// each affected chat so watching clients see every chat that reappeared.
// ErrChildUnarchiveParentArchived is returned by UnarchiveChat when a
// child unarchive is rejected because the parent is still archived.
// The patchChat handler maps this to a 400 response.
var ErrChildUnarchiveParentArchived = xerrors.New(
"cannot unarchive child chat while parent is archived",
)
// UnarchiveChat unarchives a chat family and broadcasts created events.
// Root chats cascade through UnarchiveChatByID. Child chats run under
// a row-level lock on the child (GetChatByIDForUpdate) with an
// in-transaction re-read of the parent, returning
// ErrChildUnarchiveParentArchived when the parent is archived and a
// no-op when the child is already active.
//
// The child is locked before the parent is read to avoid deadlocking
// with a concurrent ArchiveChatByID cascade, which visits child rows
// before the parent.
func (p *Server) UnarchiveChat(ctx context.Context, chat database.Chat) error {
if chat.ID == uuid.Nil {
return xerrors.New("chat_id is required")
}
return p.applyChatLifecycleTransition(
ctx,
chat.ID,
"unarchive",
codersdk.ChatWatchEventKindCreated,
p.db.UnarchiveChatByID,
)
if !chat.ParentChatID.Valid {
return p.applyChatLifecycleTransition(
ctx,
chat.ID,
"unarchive",
codersdk.ChatWatchEventKindCreated,
p.db.UnarchiveChatByID,
)
}
var updated []database.Chat
if err := p.db.InTx(func(tx database.Store) error {
locked, err := tx.GetChatByIDForUpdate(ctx, chat.ID)
if err != nil {
return xerrors.Errorf("lock child for unarchive: %w", err)
}
if !locked.Archived {
// Already unarchived by a concurrent caller; idempotent no-op.
return nil
}
parent, err := tx.GetChatByID(ctx, chat.ParentChatID.UUID)
if err != nil {
return xerrors.Errorf("load parent chat: %w", err)
}
if parent.Archived {
return ErrChildUnarchiveParentArchived
}
updated, err = tx.UnarchiveChatByID(ctx, chat.ID)
if err != nil {
return xerrors.Errorf("unarchive child chat: %w", err)
}
return nil
}, nil); err != nil {
if errors.Is(err, ErrChildUnarchiveParentArchived) {
return ErrChildUnarchiveParentArchived
}
return err
}
p.publishChatPubsubEvents(updated, codersdk.ChatWatchEventKindCreated)
return nil
}
func (p *Server) applyChatLifecycleTransition(
+137 -10
View File
@@ -644,11 +644,19 @@ func TestExploreSubagentIsReadOnly(t *testing.T) {
require.True(t, requestHasSystemSubstring(childRequests[0], "You are in Explore Mode as a delegated sub-agent."))
require.False(t, requestHasSystemSubstring(rootRequests[0], "You are in Explore Mode as a delegated sub-agent."))
allChats, err := db.GetChats(dbauthz.AsChatd(ctx), database.GetChatsParams{OwnerID: user.UserID})
rootChats, err := db.GetChats(dbauthz.AsChatd(ctx), database.GetChatsParams{OwnerID: user.UserID})
require.NoError(t, err)
rootIDs := make([]uuid.UUID, 0, len(rootChats))
for _, root := range rootChats {
rootIDs = append(rootIDs, root.Chat.ID)
}
childRows, err := db.GetChildChatsByParentIDs(dbauthz.AsChatd(ctx), database.GetChildChatsByParentIDsParams{
ParentIds: rootIDs,
})
require.NoError(t, err)
var exploreChildren []database.Chat
for _, candidate := range allChats {
if candidate.Chat.ParentChatID.Valid && candidate.Chat.Mode.Valid && candidate.Chat.Mode.ChatMode == database.ChatModeExplore {
for _, candidate := range childRows {
if candidate.Chat.Mode.Valid && candidate.Chat.Mode.ChatMode == database.ChatModeExplore {
exploreChildren = append(exploreChildren, candidate.Chat)
}
}
@@ -733,6 +741,127 @@ func TestArchiveChatMovesPendingChatToWaiting(t *testing.T) {
require.Zero(t, fromDB.PinOrder)
}
// TestUnarchiveChildChat covers the deterministic branches of the
// Server.UnarchiveChat child path: happy path, archived-parent reject,
// and already-active no-op.
func TestUnarchiveChildChat(t *testing.T) {
t.Parallel()
t.Run("ChildWithActiveParentUnarchives", func(t *testing.T) {
t.Parallel()
db, ps := dbtestutil.NewDB(t)
replica := newTestServer(t, db, ps, uuid.New())
ctx := testutil.Context(t, testutil.WaitLong)
user, org, model := seedChatDependencies(ctx, t, db)
parent, child := insertParentWithArchivedChild(ctx, t, db, user, org, model)
require.NoError(t, replica.UnarchiveChat(ctx, child))
dbChild, err := db.GetChatByID(ctx, child.ID)
require.NoError(t, err)
require.False(t, dbChild.Archived, "child should be unarchived")
dbParent, err := db.GetChatByID(ctx, parent.ID)
require.NoError(t, err)
require.False(t, dbParent.Archived, "parent should stay active")
})
t.Run("ChildWithArchivedParentRejected", func(t *testing.T) {
t.Parallel()
db, ps := dbtestutil.NewDB(t)
replica := newTestServer(t, db, ps, uuid.New())
ctx := testutil.Context(t, testutil.WaitLong)
user, org, model := seedChatDependencies(ctx, t, db)
parent, child := insertParentWithArchivedChild(ctx, t, db, user, org, model)
_, err := db.ArchiveChatByID(ctx, parent.ID)
require.NoError(t, err)
err = replica.UnarchiveChat(ctx, child)
require.ErrorIs(t, err, chatd.ErrChildUnarchiveParentArchived)
dbChild, err := db.GetChatByID(ctx, child.ID)
require.NoError(t, err)
require.True(t, dbChild.Archived, "child should remain archived")
})
t.Run("AlreadyActiveChildNoOp", func(t *testing.T) {
t.Parallel()
db, ps := dbtestutil.NewDB(t)
replica := newTestServer(t, db, ps, uuid.New())
ctx := testutil.Context(t, testutil.WaitLong)
user, org, model := seedChatDependencies(ctx, t, db)
_, child := insertParentWithActiveChild(ctx, t, db, user, org, model)
require.NoError(t, replica.UnarchiveChat(ctx, child))
dbChild, err := db.GetChatByID(ctx, child.ID)
require.NoError(t, err)
require.False(t, dbChild.Archived, "child should stay active")
})
}
// insertParentWithActiveChild creates a parent chat and an active
// child chat linked to it. Both are returned in their initial
// (active) state.
func insertParentWithActiveChild(
ctx context.Context,
t *testing.T,
db database.Store,
user database.User,
org database.Organization,
model database.ChatModelConfig,
) (parent database.Chat, child database.Chat) {
t.Helper()
var err error
parent, err = db.InsertChat(ctx, database.InsertChatParams{
OrganizationID: org.ID,
OwnerID: user.ID,
Status: database.ChatStatusWaiting,
ClientType: database.ChatClientTypeUi,
LastModelConfigID: model.ID,
Title: "parent",
})
require.NoError(t, err)
child, err = db.InsertChat(ctx, database.InsertChatParams{
OrganizationID: org.ID,
OwnerID: user.ID,
Status: database.ChatStatusWaiting,
ClientType: database.ChatClientTypeUi,
LastModelConfigID: model.ID,
Title: "child",
ParentChatID: uuid.NullUUID{UUID: parent.ID, Valid: true},
RootChatID: uuid.NullUUID{UUID: parent.ID, Valid: true},
})
require.NoError(t, err)
return parent, child
}
// insertParentWithArchivedChild creates an active parent and an
// individually-archived child. The returned child reflects its
// current (archived) state in the DB.
func insertParentWithArchivedChild(
ctx context.Context,
t *testing.T,
db database.Store,
user database.User,
org database.Organization,
model database.ChatModelConfig,
) (parent database.Chat, child database.Chat) {
t.Helper()
parent, child = insertParentWithActiveChild(ctx, t, db, user, org, model)
_, err := db.ArchiveChatByID(ctx, child.ID)
require.NoError(t, err)
child, err = db.GetChatByID(ctx, child.ID)
require.NoError(t, err)
return parent, child
}
func TestArchiveChatInterruptsActiveProcessing(t *testing.T) {
t.Parallel()
@@ -4976,15 +5105,13 @@ func TestComputerUseSubagentToolsAndModel(t *testing.T) {
// 6. Verify the child chat has Mode = computer_use in
// the DB.
allChats, err := db.GetChats(ctx, database.GetChatsParams{
OwnerID: user.ID,
childRows, err := db.GetChildChatsByParentIDs(ctx, database.GetChildChatsByParentIDsParams{
ParentIds: []uuid.UUID{chat.ID},
})
require.NoError(t, err)
var children []database.Chat
for _, c := range allChats {
if c.Chat.ParentChatID.Valid && c.Chat.ParentChatID.UUID == chat.ID {
children = append(children, c.Chat)
}
children := make([]database.Chat, 0, len(childRows))
for _, row := range childRows {
children = append(children, row.Chat)
}
require.Len(t, children, 1)
require.True(t, children[0].Mode.Valid)
+9 -29
View File
@@ -66,9 +66,9 @@ type Store interface {
UpsertChatDiffStatusReference(
ctx context.Context, arg database.UpsertChatDiffStatusReferenceParams,
) (database.ChatDiffStatus, error)
GetChats(
ctx context.Context, arg database.GetChatsParams,
) ([]database.GetChatsRow, error)
GetChatsByWorkspaceIDs(
ctx context.Context, ids []uuid.UUID,
) ([]database.Chat, error)
}
// EventPublisher notifies the frontend of diff status changes.
@@ -277,7 +277,6 @@ func (w *Worker) tick(ctx context.Context) {
// MarkStaleParams holds the arguments for Worker.MarkStale.
type MarkStaleParams struct {
WorkspaceID uuid.UUID
OwnerID uuid.UUID
Branch string
Origin string
// ChatID, when set, targets a single chat instead of
@@ -306,9 +305,11 @@ func (w *Worker) MarkStale(ctx context.Context, p MarkStaleParams) {
return
}
chatRows, err := w.store.GetChats(ctx, database.GetChatsParams{
OwnerID: p.OwnerID,
})
// Broadcast path: scope by workspace. GetChatsByWorkspaceIDs
// filters archived=false, which is intentional: archived
// chats aren't in the active sidebar and don't need refreshed
// git refs.
chats, err := w.store.GetChatsByWorkspaceIDs(ctx, []uuid.UUID{p.WorkspaceID})
if err != nil {
w.logger.Warn(ctx, "list chats for git ref storage",
slog.F("workspace_id", p.WorkspaceID),
@@ -316,12 +317,7 @@ func (w *Worker) MarkStale(ctx context.Context, p MarkStaleParams) {
return
}
chats := make([]database.Chat, len(chatRows))
for i, row := range chatRows {
chats[i] = row.Chat
}
for _, chat := range filterChatsByWorkspaceID(chats, p.WorkspaceID) {
for _, chat := range chats {
w.markStaleSingle(ctx, chat.ID, p.Branch, p.Origin)
}
}
@@ -403,19 +399,3 @@ func (w *Worker) RefreshChat(
return &upserted, nil
}
// filterChatsByWorkspaceID returns only chats associated with
// the given workspace.
func filterChatsByWorkspaceID(
chats []database.Chat,
workspaceID uuid.UUID,
) []database.Chat {
filtered := make([]database.Chat, 0, len(chats))
for _, chat := range chats {
if !chat.WorkspaceID.Valid || chat.WorkspaceID.UUID != workspaceID {
continue
}
filtered = append(filtered, chat)
}
return filtered
}
+23 -36
View File
@@ -606,7 +606,6 @@ func TestWorker_MarkStale_UpsertAndPublish(t *testing.T) {
ownerID := uuid.New()
chat1 := uuid.New()
chat2 := uuid.New()
chatOther := uuid.New()
var mu sync.Mutex
var upsertRefCalls []database.UpsertChatDiffStatusReferenceParams
@@ -615,13 +614,12 @@ func TestWorker_MarkStale_UpsertAndPublish(t *testing.T) {
ctrl := gomock.NewController(t)
store := dbmock.NewMockStore(ctrl)
store.EXPECT().GetChats(gomock.Any(), gomock.Any()).
DoAndReturn(func(_ context.Context, arg database.GetChatsParams) ([]database.GetChatsRow, error) {
require.Equal(t, ownerID, arg.OwnerID)
return []database.GetChatsRow{
{Chat: database.Chat{ID: chat1, OwnerID: ownerID, WorkspaceID: uuid.NullUUID{UUID: workspaceID, Valid: true}}},
{Chat: database.Chat{ID: chat2, OwnerID: ownerID, WorkspaceID: uuid.NullUUID{UUID: workspaceID, Valid: true}}},
{Chat: database.Chat{ID: chatOther, OwnerID: ownerID, WorkspaceID: uuid.NullUUID{UUID: uuid.New(), Valid: true}}},
store.EXPECT().GetChatsByWorkspaceIDs(gomock.Any(), gomock.Any()).
DoAndReturn(func(_ context.Context, ids []uuid.UUID) ([]database.Chat, error) {
require.Equal(t, []uuid.UUID{workspaceID}, ids)
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}},
}, nil
})
store.EXPECT().UpsertChatDiffStatusReference(gomock.Any(), gomock.Any()).DoAndReturn(func(_ context.Context, arg database.UpsertChatDiffStatusReferenceParams) (database.ChatDiffStatus, error) {
@@ -646,7 +644,6 @@ func TestWorker_MarkStale_UpsertAndPublish(t *testing.T) {
worker.MarkStale(ctx, gitsync.MarkStaleParams{
WorkspaceID: workspaceID,
OwnerID: ownerID,
Branch: "feature",
Origin: "https://github.com/owner/repo",
})
@@ -672,16 +669,12 @@ func TestWorker_MarkStale_NoMatchingChats(t *testing.T) {
ctx := testutil.Context(t, testutil.WaitShort)
workspaceID := uuid.New()
ownerID := uuid.New()
ctrl := gomock.NewController(t)
store := dbmock.NewMockStore(ctrl)
store.EXPECT().GetChats(gomock.Any(), gomock.Any()).
Return([]database.GetChatsRow{
{Chat: database.Chat{ID: uuid.New(), OwnerID: ownerID, WorkspaceID: uuid.NullUUID{UUID: uuid.New(), Valid: true}}},
{Chat: database.Chat{ID: uuid.New(), OwnerID: ownerID, WorkspaceID: uuid.NullUUID{UUID: uuid.New(), Valid: true}}},
}, nil)
store.EXPECT().GetChatsByWorkspaceIDs(gomock.Any(), gomock.Any()).
Return(nil, nil)
mClock := quartz.NewMock(t)
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
@@ -690,7 +683,6 @@ func TestWorker_MarkStale_NoMatchingChats(t *testing.T) {
worker.MarkStale(ctx, gitsync.MarkStaleParams{
WorkspaceID: workspaceID,
OwnerID: ownerID,
Branch: "main",
Origin: "https://github.com/x/y",
})
@@ -710,10 +702,10 @@ func TestWorker_MarkStale_UpsertFails_ContinuesNext(t *testing.T) {
ctrl := gomock.NewController(t)
store := dbmock.NewMockStore(ctrl)
store.EXPECT().GetChats(gomock.Any(), gomock.Any()).
Return([]database.GetChatsRow{
{Chat: database.Chat{ID: chat1, OwnerID: ownerID, WorkspaceID: uuid.NullUUID{UUID: workspaceID, Valid: true}}},
{Chat: database.Chat{ID: chat2, OwnerID: ownerID, WorkspaceID: uuid.NullUUID{UUID: workspaceID, Valid: true}}},
store.EXPECT().GetChatsByWorkspaceIDs(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}},
}, nil)
store.EXPECT().UpsertChatDiffStatusReference(gomock.Any(), gomock.Any()).
DoAndReturn(func(_ context.Context, arg database.UpsertChatDiffStatusReferenceParams) (database.ChatDiffStatus, error) {
@@ -735,7 +727,6 @@ func TestWorker_MarkStale_UpsertFails_ContinuesNext(t *testing.T) {
worker.MarkStale(ctx, gitsync.MarkStaleParams{
WorkspaceID: workspaceID,
OwnerID: ownerID,
Branch: "dev",
Origin: "https://github.com/a/b",
})
@@ -743,14 +734,14 @@ func TestWorker_MarkStale_UpsertFails_ContinuesNext(t *testing.T) {
assert.Equal(t, int32(1), publishCount.Load())
}
func TestWorker_MarkStale_GetChatsFails(t *testing.T) {
func TestWorker_MarkStale_GetChatsByWorkspaceIDsFails(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitShort)
ctrl := gomock.NewController(t)
store := dbmock.NewMockStore(ctrl)
store.EXPECT().GetChats(gomock.Any(), gomock.Any()).
store.EXPECT().GetChatsByWorkspaceIDs(gomock.Any(), gomock.Any()).
Return(nil, fmt.Errorf("db error"))
mClock := quartz.NewMock(t)
@@ -760,7 +751,6 @@ func TestWorker_MarkStale_GetChatsFails(t *testing.T) {
worker.MarkStale(ctx, gitsync.MarkStaleParams{
WorkspaceID: uuid.New(),
OwnerID: uuid.New(),
Branch: "main",
Origin: "https://github.com/x/y",
})
@@ -817,7 +807,6 @@ func TestWorker_MarkStale_EmptyBranchOrOrigin(t *testing.T) {
worker.MarkStale(ctx, gitsync.MarkStaleParams{
WorkspaceID: uuid.New(),
OwnerID: uuid.New(),
Branch: tc.branch,
Origin: tc.origin,
})
@@ -838,8 +827,8 @@ func TestWorker_MarkStale_WithChatID(t *testing.T) {
ctrl := gomock.NewController(t)
store := dbmock.NewMockStore(ctrl)
// GetChats should NOT be called when a specific chat ID is provided.
store.EXPECT().GetChats(gomock.Any(), gomock.Any()).Times(0)
// GetChatsByWorkspaceIDs should NOT be called when a specific chat ID is provided.
store.EXPECT().GetChatsByWorkspaceIDs(gomock.Any(), gomock.Any()).Times(0)
store.EXPECT().UpsertChatDiffStatusReference(gomock.Any(), gomock.Any()).DoAndReturn(func(_ context.Context, arg database.UpsertChatDiffStatusReferenceParams) (database.ChatDiffStatus, error) {
mu.Lock()
upsertRefCalls = append(upsertRefCalls, arg)
@@ -862,7 +851,6 @@ func TestWorker_MarkStale_WithChatID(t *testing.T) {
worker.MarkStale(ctx, gitsync.MarkStaleParams{
WorkspaceID: uuid.New(),
OwnerID: uuid.New(),
Branch: "my-branch",
Origin: "https://github.com/org/repo",
ChatID: targetChat,
@@ -897,13 +885,13 @@ func TestWorker_MarkStale_NilChatID_Broadcasts(t *testing.T) {
ctrl := gomock.NewController(t)
store := dbmock.NewMockStore(ctrl)
// GetChats IS called because a nil ChatID triggers the
// workspace-wide broadcast path.
store.EXPECT().GetChats(gomock.Any(), gomock.Any()).
DoAndReturn(func(_ context.Context, arg database.GetChatsParams) ([]database.GetChatsRow, error) {
require.Equal(t, ownerID, arg.OwnerID)
return []database.GetChatsRow{
{Chat: database.Chat{ID: chat1, OwnerID: ownerID, WorkspaceID: uuid.NullUUID{UUID: workspaceID, Valid: true}}},
// Broadcast path: GetChatsByWorkspaceIDs scopes the query to
// the workspace directly; no post-filtering needed.
store.EXPECT().GetChatsByWorkspaceIDs(gomock.Any(), gomock.Any()).
DoAndReturn(func(_ context.Context, ids []uuid.UUID) ([]database.Chat, error) {
require.Equal(t, []uuid.UUID{workspaceID}, ids)
return []database.Chat{
{ID: chat1, OwnerID: ownerID, WorkspaceID: uuid.NullUUID{UUID: workspaceID, Valid: true}},
}, nil
})
store.EXPECT().UpsertChatDiffStatusReference(gomock.Any(), gomock.Any()).DoAndReturn(func(_ context.Context, arg database.UpsertChatDiffStatusReferenceParams) (database.ChatDiffStatus, error) {
@@ -928,7 +916,6 @@ func TestWorker_MarkStale_NilChatID_Broadcasts(t *testing.T) {
// Zero-value ChatID (uuid.Nil) triggers broadcast.
worker.MarkStale(ctx, gitsync.MarkStaleParams{
WorkspaceID: workspaceID,
OwnerID: ownerID,
Branch: "main",
Origin: "https://github.com/org/repo",
})