mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
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:
+58
-9
@@ -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
@@ -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)
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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",
|
||||
})
|
||||
|
||||
Reference in New Issue
Block a user