From 79c6547afba0cf2b016ed9c6432ffccd6eb2778c Mon Sep 17 00:00:00 2001 From: Lifei Zhou Date: Wed, 24 Jun 2026 23:15:55 +1000 Subject: [PATCH] feat (ui): Use ACP steer to send the message in queue in ACP path (#9957) --- .../acp/__tests__/chatNotifications.test.ts | 4 + .../__tests__/chatSessionController.test.ts | 184 +++++++++- .../acp/__tests__/chatSessionStore.test.ts | 181 +++++++++ .../sessionNotificationAdapter.test.ts | 147 +++++++- ui/desktop/src/acp/adapter/messages.ts | 54 ++- ui/desktop/src/acp/adapter/shared.ts | 25 +- ui/desktop/src/acp/chatSessionController.ts | 31 +- ui/desktop/src/acp/chatSessionStore.ts | 140 ++++++- ui/desktop/src/acp/prompt.ts | 14 + .../src/acp/sessionNotificationAdapter.ts | 20 +- ui/desktop/src/acpChatFeatureFlag.ts | 2 +- ui/desktop/src/components/BaseChat.tsx | 4 + ui/desktop/src/components/ChatInput.tsx | 109 +++++- ui/desktop/src/components/MessageQueue.tsx | 346 ++++++++++-------- ui/desktop/src/hooks/useAcpChatSession.ts | 66 +++- ui/desktop/src/hooks/useChatSessionTypes.ts | 2 + ui/desktop/src/hooks/useChatStream.ts | 1 + ui/desktop/src/i18n/messages/en.json | 3 + ui/desktop/src/i18n/messages/es.json | 3 + ui/desktop/src/i18n/messages/hi.json | 3 + ui/desktop/src/i18n/messages/ja.json | 3 + ui/desktop/src/i18n/messages/ko.json | 3 + ui/desktop/src/i18n/messages/ru.json | 3 + ui/desktop/src/i18n/messages/tr.json | 3 + ui/desktop/src/i18n/messages/zh-CN.json | 3 + 25 files changed, 1175 insertions(+), 179 deletions(-) diff --git a/ui/desktop/src/acp/__tests__/chatNotifications.test.ts b/ui/desktop/src/acp/__tests__/chatNotifications.test.ts index 440209c93..98393a60d 100644 --- a/ui/desktop/src/acp/__tests__/chatNotifications.test.ts +++ b/ui/desktop/src/acp/__tests__/chatNotifications.test.ts @@ -62,6 +62,8 @@ function snapshotWithName(name: string): AcpChatSessionSnapshot { chatState: ChatState.Idle, sessionLoadError: undefined, activePromptAttemptId: null, + activeRunId: null, + pendingCancelPromptAttemptId: null, }; } @@ -81,6 +83,8 @@ function snapshotWithoutSession(): AcpChatSessionSnapshot { chatState: ChatState.Idle, sessionLoadError: undefined, activePromptAttemptId: null, + activeRunId: null, + pendingCancelPromptAttemptId: null, }; } diff --git a/ui/desktop/src/acp/__tests__/chatSessionController.test.ts b/ui/desktop/src/acp/__tests__/chatSessionController.test.ts index cb598fa2e..be0f0b9a7 100644 --- a/ui/desktop/src/acp/__tests__/chatSessionController.test.ts +++ b/ui/desktop/src/acp/__tests__/chatSessionController.test.ts @@ -1,8 +1,19 @@ import { beforeEach, describe, expect, it, vi } from 'vitest'; -import type { Session } from '../../api'; +import type { Message, Session } from '../../api'; +import { ChatState } from '../../types/chatState'; import { acpChatSessionController } from '../chatSessionController'; -import { acpChatSessionActions, acpChatSessionStore } from '../chatSessionStore'; -import { acpLoadSession, isAcpSessionLoadInFlight, sessionInfoToSession } from '../sessions'; +import { + acpChatSessionActions, + acpChatSessionStore, + type AcpChatSessionSnapshot, +} from '../chatSessionStore'; +import { acpCancelPrompt, acpPromptSession } from '../prompt'; +import { + acpLoadSession, + acpTruncateSessionConversation, + isAcpSessionLoadInFlight, + sessionInfoToSession, +} from '../sessions'; vi.mock('../../utils/extensionErrorUtils', () => ({ showExtensionLoadResults: vi.fn(), @@ -20,7 +31,10 @@ vi.mock('../chatSessionStore', () => ({ finishPromptAttemptIfCurrent: vi.fn(), isCurrentPromptAttempt: vi.fn(), setMessages: vi.fn(), + addPendingLocalSteerMessage: vi.fn(), clearActivePromptAttempt: vi.fn(), + startPromptCancellation: vi.fn(), + clearPromptCancellation: vi.fn(), setChatState: vi.fn(), setSessionMetadata: vi.fn(), setSessionLoadError: vi.fn(), @@ -35,8 +49,23 @@ vi.mock('../sessions', () => ({ acpTruncateSessionConversation: vi.fn(), })); +vi.mock('../prompt', () => ({ + acpCancelPrompt: vi.fn(), + acpPromptSession: vi.fn(), +})); + const SESSION_ID = 'session-1'; +function userMessage(): Message & { id: string } { + return { + id: 'message-1', + role: 'user', + created: 123, + content: [{ type: 'text', text: 'Hello' }], + metadata: { userVisible: true, agentVisible: true }, + }; +} + function loadedSession(): Session { return { id: SESSION_ID, @@ -63,6 +92,27 @@ function mockLoadResult() { } as Awaited>; } +function snapshotWithActivePrompt(activePromptAttemptId: string | null): AcpChatSessionSnapshot { + return { + session: undefined, + messages: [], + tokenState: { + inputTokens: 0, + outputTokens: 0, + totalTokens: 0, + accumulatedInputTokens: 0, + accumulatedOutputTokens: 0, + accumulatedTotalTokens: 0, + }, + notifications: [], + chatState: activePromptAttemptId ? ChatState.Streaming : ChatState.Idle, + sessionLoadError: undefined, + activePromptAttemptId, + activeRunId: activePromptAttemptId ? 'run-1' : null, + pendingCancelPromptAttemptId: null, + }; +} + describe('acpChatSessionController.loadSession', () => { beforeEach(() => { vi.clearAllMocks(); @@ -97,3 +147,131 @@ describe('acpChatSessionController.loadSession', () => { ); }); }); + +describe('acpChatSessionController.stop', () => { + beforeEach(() => { + vi.clearAllMocks(); + vi.mocked(acpCancelPrompt).mockResolvedValue(undefined); + }); + + it('marks cancellation pending while clearing visible prompt activity', () => { + vi.mocked(acpChatSessionStore.getSnapshot).mockReturnValue( + snapshotWithActivePrompt('attempt-1') + ); + + acpChatSessionController.stop(SESSION_ID); + + expect(acpChatSessionActions.startPromptCancellation).toHaveBeenCalledWith( + SESSION_ID, + 'attempt-1' + ); + expect(acpCancelPrompt).toHaveBeenCalledWith(SESSION_ID); + }); +}); + +describe('acpChatSessionController.submitMessage', () => { + beforeEach(() => { + vi.clearAllMocks(); + vi.mocked(acpChatSessionStore.getSnapshot).mockReturnValue(snapshotWithActivePrompt(null)); + vi.mocked(acpPromptSession).mockResolvedValue({ stopReason: 'cancelled' } as never); + vi.mocked(acpChatSessionActions.clearPromptCancellation).mockReturnValue(undefined); + vi.mocked(acpChatSessionActions.finishPromptAttemptIfCurrent).mockReturnValue(true); + }); + + it('clears a pending cancellation barrier when the original prompt settles', async () => { + vi.mocked(acpChatSessionActions.clearPromptCancellation).mockReturnValueOnce( + snapshotWithActivePrompt(null) + ); + const onFinish = vi.fn(); + + await acpChatSessionController.submitMessage(SESSION_ID, userMessage(), { + getCurrentSnapshot: () => snapshotWithActivePrompt(null), + onFinish, + }); + + expect(acpChatSessionActions.clearPromptCancellation).toHaveBeenCalledWith( + SESSION_ID, + expect.any(String) + ); + expect(acpChatSessionActions.finishPromptAttemptIfCurrent).not.toHaveBeenCalled(); + expect(onFinish).not.toHaveBeenCalled(); + }); + + it('rejects while a cancellation barrier is pending', async () => { + vi.mocked(acpChatSessionStore.getSnapshot).mockReturnValue({ + ...snapshotWithActivePrompt(null), + pendingCancelPromptAttemptId: 'attempt-1', + }); + + await expect( + acpChatSessionController.submitMessage(SESSION_ID, userMessage(), { + getCurrentSnapshot: () => snapshotWithActivePrompt(null), + onFinish: vi.fn(), + }) + ).rejects.toThrow('Cannot submit while prompt cancellation is pending'); + + expect(acpChatSessionActions.startPromptAttempt).not.toHaveBeenCalled(); + expect(acpPromptSession).not.toHaveBeenCalled(); + }); +}); + +describe('acpChatSessionController.updateMessage', () => { + beforeEach(() => { + vi.clearAllMocks(); + vi.mocked(acpTruncateSessionConversation).mockResolvedValue(undefined as never); + vi.mocked(acpChatSessionStore.getSnapshot).mockReturnValue(snapshotWithActivePrompt(null)); + }); + + it('rejects edits before truncating while cancellation is pending', async () => { + vi.mocked(acpChatSessionStore.getSnapshot).mockReturnValue({ + ...snapshotWithActivePrompt(null), + pendingCancelPromptAttemptId: 'attempt-1', + }); + const existingMessage = userMessage(); + const currentSnapshot: AcpChatSessionSnapshot = { + ...snapshotWithActivePrompt(null), + messages: [existingMessage], + }; + + await expect( + acpChatSessionController.updateMessage(SESSION_ID, existingMessage.id, 'Updated', 'edit', { + getCurrentSnapshot: () => currentSnapshot, + onFinish: vi.fn(), + }) + ).rejects.toThrow('Cannot submit while prompt cancellation is pending'); + + expect(acpChatSessionActions.setChatState).not.toHaveBeenCalledWith( + SESSION_ID, + ChatState.Thinking + ); + expect(acpTruncateSessionConversation).not.toHaveBeenCalled(); + expect(acpChatSessionActions.setMessages).not.toHaveBeenCalled(); + expect(acpPromptSession).not.toHaveBeenCalled(); + }); + + it('rejects edits before truncating while a prompt is active', async () => { + vi.mocked(acpChatSessionStore.getSnapshot).mockReturnValue( + snapshotWithActivePrompt('attempt-1') + ); + const existingMessage = userMessage(); + const currentSnapshot: AcpChatSessionSnapshot = { + ...snapshotWithActivePrompt('attempt-1'), + messages: [existingMessage], + }; + + await expect( + acpChatSessionController.updateMessage(SESSION_ID, existingMessage.id, 'Updated', 'edit', { + getCurrentSnapshot: () => currentSnapshot, + onFinish: vi.fn(), + }) + ).rejects.toThrow('Cannot update message while prompt is active'); + + expect(acpChatSessionActions.setChatState).not.toHaveBeenCalledWith( + SESSION_ID, + ChatState.Thinking + ); + expect(acpTruncateSessionConversation).not.toHaveBeenCalled(); + expect(acpChatSessionActions.setMessages).not.toHaveBeenCalled(); + expect(acpPromptSession).not.toHaveBeenCalled(); + }); +}); diff --git a/ui/desktop/src/acp/__tests__/chatSessionStore.test.ts b/ui/desktop/src/acp/__tests__/chatSessionStore.test.ts index feb750fdc..23834d0b9 100644 --- a/ui/desktop/src/acp/__tests__/chatSessionStore.test.ts +++ b/ui/desktop/src/acp/__tests__/chatSessionStore.test.ts @@ -129,6 +129,44 @@ function agentMessageChunkNotification( }; } +function userSteerChunkNotification( + sessionId: string, + messageId: string, + text: string +): SessionNotification { + return { + sessionId, + update: { + sessionUpdate: 'user_message_chunk', + messageId, + content: { + type: 'text', + text, + }, + _meta: { + goose: { + messageId, + steer: true, + }, + }, + } as SessionNotification['update'], + }; +} + +function activeRunNotification(sessionId: string, activeRunId: string | null): SessionNotification { + return { + sessionId, + update: { + sessionUpdate: 'session_info_update', + _meta: { + goose: { + activeRunId, + }, + }, + } as SessionNotification['update'], + }; +} + describe('acpChatSessionStore', () => { const sessionIds = new Set(); const sessionId = (id: string): string => { @@ -224,6 +262,149 @@ describe('acpChatSessionStore', () => { expect(snapshot.chatState).toBe(ChatState.Streaming); }); + it('tracks prompt cancellation separately from visible prompt activity', () => { + const currentSessionId = sessionId('session-1'); + + acpChatSessionActions.startPromptAttempt(currentSessionId, 'attempt-1'); + const cancellationSnapshot = acpChatSessionActions.startPromptCancellation( + currentSessionId, + 'attempt-1' + ); + + expect(cancellationSnapshot).toMatchObject({ + activePromptAttemptId: null, + pendingCancelPromptAttemptId: 'attempt-1', + chatState: ChatState.Idle, + }); + + const staleClearSnapshot = acpChatSessionActions.clearPromptCancellation( + currentSessionId, + 'attempt-2' + ); + expect(staleClearSnapshot).toBeUndefined(); + expect(acpChatSessionStore.getSnapshot(currentSessionId)?.pendingCancelPromptAttemptId).toBe( + 'attempt-1' + ); + + const clearedSnapshot = acpChatSessionActions.clearPromptCancellation( + currentSessionId, + 'attempt-1' + ); + + expect(clearedSnapshot?.pendingCancelPromptAttemptId).toBeNull(); + }); + + it('removes pending local steer messages when cancellation starts', () => { + const currentSessionId = sessionId('session-1'); + const localSteerMessage = { + ...message('steer-1', 'hello'), + metadata: { userVisible: true, agentVisible: true, steer: true }, + }; + + acpChatSessionActions.startPromptAttempt(currentSessionId, 'attempt-1'); + acpChatSessionActions.addPendingLocalSteerMessage(currentSessionId, localSteerMessage); + + expect(acpChatSessionStore.getSnapshot(currentSessionId)?.messages).toHaveLength(1); + + const cancellationSnapshot = acpChatSessionActions.startPromptCancellation( + currentSessionId, + 'attempt-1' + ); + + expect(cancellationSnapshot?.messages).toEqual([]); + }); + + it('keeps confirmed local steer messages when cancellation starts', () => { + const currentSessionId = sessionId('session-1'); + const localSteerMessage = { + ...message('steer-1', 'hello'), + metadata: { userVisible: true, agentVisible: true, steer: true }, + }; + + acpChatSessionActions.startPromptAttempt(currentSessionId, 'attempt-1'); + acpChatSessionActions.addPendingLocalSteerMessage(currentSessionId, localSteerMessage); + acpChatSessionActions.applyAcpSessionNotification( + userSteerChunkNotification(currentSessionId, 'steer-1', 'hello') + ); + + const cancellationSnapshot = acpChatSessionActions.startPromptCancellation( + currentSessionId, + 'attempt-1' + ); + + expect(cancellationSnapshot?.messages).toHaveLength(1); + expect(cancellationSnapshot?.messages[0].id).toBe('steer-1'); + }); + + it('preserves steer text accumulation when another local steer is added', () => { + const currentSessionId = sessionId('session-1'); + const firstSteerMessage = { + ...message('steer-1', 'hello'), + metadata: { userVisible: true, agentVisible: true, steer: true }, + }; + const secondSteerMessage = { + ...message('steer-2', 'second'), + metadata: { userVisible: true, agentVisible: true, steer: true }, + }; + + acpChatSessionActions.startPromptAttempt(currentSessionId, 'attempt-1'); + acpChatSessionActions.addPendingLocalSteerMessage(currentSessionId, firstSteerMessage); + acpChatSessionActions.applyAcpSessionNotification( + userSteerChunkNotification(currentSessionId, 'steer-1', 'hel') + ); + + acpChatSessionActions.addPendingLocalSteerMessage(currentSessionId, secondSteerMessage); + const snapshot = acpChatSessionActions.applyAcpSessionNotification( + userSteerChunkNotification(currentSessionId, 'steer-1', 'lo') + ); + + const firstMessage = snapshot.messages.find((item) => item.id === 'steer-1'); + expect(firstMessage?.content[0]).toMatchObject({ type: 'text', text: 'hello' }); + }); + + it('stores active run ids from session info notifications', () => { + const currentSessionId = sessionId('session-1'); + + const snapshot = acpChatSessionActions.applyAcpSessionNotification( + activeRunNotification(currentSessionId, 'run-1') + ); + + expect(snapshot.activeRunId).toBe('run-1'); + expect(acpChatSessionStore.getSnapshot(currentSessionId)?.activeRunId).toBe('run-1'); + + const clearedSnapshot = acpChatSessionActions.applyAcpSessionNotification( + activeRunNotification(currentSessionId, null) + ); + + expect(clearedSnapshot.activeRunId).toBeNull(); + }); + + it('clears active run ids when the prompt attempt finishes', () => { + const currentSessionId = sessionId('session-1'); + + acpChatSessionActions.startPromptAttempt(currentSessionId, 'attempt-1'); + acpChatSessionActions.applyAcpSessionNotification( + activeRunNotification(currentSessionId, 'run-1') + ); + + expect(acpChatSessionActions.finishPromptAttemptIfCurrent(currentSessionId, 'attempt-1')).toBe( + true + ); + expect(acpChatSessionStore.getSnapshot(currentSessionId)?.activeRunId).toBeNull(); + }); + + it('clears active run ids before replaying a session load', () => { + const currentSessionId = sessionId('session-1'); + + acpChatSessionActions.applyAcpSessionNotification( + activeRunNotification(currentSessionId, 'run-1') + ); + + const snapshot = acpChatSessionActions.startSessionLoad(currentSessionId); + + expect(snapshot.activeRunId).toBeNull(); + }); + it('stores ACP tool notifications and clears them for a new prompt attempt', () => { const currentSessionId = sessionId('session-1'); diff --git a/ui/desktop/src/acp/__tests__/sessionNotificationAdapter.test.ts b/ui/desktop/src/acp/__tests__/sessionNotificationAdapter.test.ts index 9a26e45b8..098839155 100644 --- a/ui/desktop/src/acp/__tests__/sessionNotificationAdapter.test.ts +++ b/ui/desktop/src/acp/__tests__/sessionNotificationAdapter.test.ts @@ -67,9 +67,24 @@ function expectOnlyMessagesChange(chatStateChanges: AcpChatStateChange[]): Messa return chatStateChange.messages; } -function expectOnlyNotificationChange( - chatStateChanges: AcpChatStateChange[] -): NotificationEvent { +function expectMessagesAndLocalSteerConfirmation( + chatStateChanges: AcpChatStateChange[], + messageId: string +): Message[] { + expect(chatStateChanges).toHaveLength(2); + + const [messagesChange, confirmationChange] = chatStateChanges; + expect(messagesChange.type).toBe('messages'); + expect(confirmationChange).toEqual({ type: 'localSteerConfirmed', messageId }); + + if (messagesChange.type !== 'messages') { + throw new Error('expected messages state change'); + } + + return messagesChange.messages; +} + +function expectOnlyNotificationChange(chatStateChanges: AcpChatStateChange[]): NotificationEvent { expect(chatStateChanges).toHaveLength(1); const [chatStateChange] = chatStateChanges; @@ -121,6 +136,132 @@ describe('createAcpSessionNotificationAdapter', () => { expect(firstContent(messages[0])).toMatchObject({ type: 'text', text: 'Hell' }); }); + it('reconciles locally rendered steer text with server chunks', () => { + const adapter = createAcpSessionNotificationAdapter([ + { + id: 'steer-1', + role: 'user', + created: 123, + content: [ + { type: 'text', text: 'hello' }, + { type: 'image', data: 'base64-image', mimeType: 'image/png' }, + ], + metadata: { userVisible: true, agentVisible: true, steer: true }, + }, + ]); + + let messages = expectMessagesAndLocalSteerConfirmation( + adapter.apply( + acpUpdate({ + sessionUpdate: 'user_message_chunk', + content: { type: 'text', text: 'hel' }, + _meta: { + goose: { + messageId: 'steer-1', + steer: true, + }, + }, + } as SessionNotification['update']) + ), + 'steer-1' + ); + + expect(firstContent(messages[0])).toMatchObject({ type: 'text', text: 'hel' }); + expect(messages[0].content[1]).toMatchObject({ + type: 'image', + data: 'base64-image', + mimeType: 'image/png', + }); + expect(messages[0].metadata.steer).toBe(true); + + messages = expectMessagesAndLocalSteerConfirmation( + adapter.apply( + acpUpdate({ + sessionUpdate: 'user_message_chunk', + content: { type: 'text', text: 'lo' }, + _meta: { + goose: { + messageId: 'steer-1', + steer: true, + }, + }, + } as SessionNotification['update']) + ), + 'steer-1' + ); + + expect(firstContent(messages[0])).toMatchObject({ type: 'text', text: 'hello' }); + + messages = expectMessagesAndLocalSteerConfirmation( + adapter.apply( + acpUpdate({ + sessionUpdate: 'user_message_chunk', + content: { type: 'image', data: 'base64-image', mimeType: 'image/png' }, + _meta: { + goose: { + messageId: 'steer-1', + steer: true, + }, + }, + } as SessionNotification['update']) + ), + 'steer-1' + ); + + expect(messages[0].content).toEqual([ + { type: 'text', text: 'hello' }, + { type: 'image', data: 'base64-image', mimeType: 'image/png' }, + ]); + }); + + it('appends repeated local steer text deltas without collapsing them', () => { + const adapter = createAcpSessionNotificationAdapter([ + { + id: 'steer-1', + role: 'user', + created: 123, + content: [{ type: 'text', text: 'haha' }], + metadata: { userVisible: true, agentVisible: true, steer: true }, + }, + ]); + + let messages = expectMessagesAndLocalSteerConfirmation( + adapter.apply( + acpUpdate({ + sessionUpdate: 'user_message_chunk', + content: { type: 'text', text: 'ha' }, + _meta: { + goose: { + messageId: 'steer-1', + steer: true, + }, + }, + } as SessionNotification['update']) + ), + 'steer-1' + ); + + expect(firstContent(messages[0])).toMatchObject({ type: 'text', text: 'ha' }); + + messages = expectMessagesAndLocalSteerConfirmation( + adapter.apply( + acpUpdate({ + sessionUpdate: 'user_message_chunk', + content: { type: 'text', text: 'ha' }, + _meta: { + goose: { + messageId: 'steer-1', + steer: true, + }, + }, + } as SessionNotification['update']) + ), + 'steer-1' + ); + + expect(firstContent(messages[0])).toMatchObject({ type: 'text', text: 'haha' }); + }); + it('maps image and thinking chunks to existing message content shapes', () => { const imageAdapter = createAcpSessionNotificationAdapter(); diff --git a/ui/desktop/src/acp/adapter/messages.ts b/ui/desktop/src/acp/adapter/messages.ts index 1b3b3caf0..9a287e83b 100644 --- a/ui/desktop/src/acp/adapter/messages.ts +++ b/ui/desktop/src/acp/adapter/messages.ts @@ -30,20 +30,29 @@ export function applyContentChunk( if (existing) { const lastContent = existing.content[existing.content.length - 1]; + if (reconcileLocalSteerTextChunk(state, existing, content, gooseMeta.steer)) { + return messagesChangeWithLocalSteerConfirmation(state, existing, gooseMeta.steer); + } + if (lastContent?.type === 'text' && content.type === 'text') { lastContent.text += content.text; } else if (content.type === 'image' && hasImageContent(existing, content)) { - return messagesChange(state); + return messagesChangeWithLocalSteerConfirmation(state, existing, gooseMeta.steer); } else { existing.content.push(content); } + + return messagesChangeWithLocalSteerConfirmation(state, existing, gooseMeta.steer); } else { state.messages.push({ ...(messageId ? { id: messageId } : {}), role, created: gooseMeta.created ?? Math.floor(Date.now() / 1000), content: [content], - metadata: { ...DEFAULT_VISIBLE_MESSAGE_METADATA }, + metadata: { + ...DEFAULT_VISIBLE_MESSAGE_METADATA, + ...(gooseMeta.steer ? { steer: true } : {}), + }, }); } @@ -148,3 +157,44 @@ function hasImageContent(message: Message, image: Extract } - | { type: 'sessionInfo'; name?: string } + | { + type: 'sessionInfo'; + name?: string; + activeRunId?: string | null; + } + | { type: 'localSteerConfirmed'; messageId: string } | { type: 'notification'; notification: NotificationEvent }; export interface AdapterState { messages: Message[]; + localSteerTextByMessageId: Map; } export interface GooseMessageMeta { messageId?: string; created?: number; + steer?: boolean; } export interface ToolIdentity { @@ -52,9 +59,25 @@ export function getGooseMessageMeta(update: { _meta?: unknown }): GooseMessageMe return { created: typeof goose.created === 'number' ? goose.created : undefined, messageId: typeof goose.messageId === 'string' ? goose.messageId : undefined, + steer: goose.steer === true ? true : undefined, }; } +export function getGooseActiveRunId(update: { _meta?: unknown }): string | null | undefined { + if (!isRecord(update._meta)) { + return undefined; + } + + const goose = update._meta.goose; + if (!isRecord(goose) || !('activeRunId' in goose)) { + return undefined; + } + + return typeof goose.activeRunId === 'string' || goose.activeRunId === null + ? goose.activeRunId + : undefined; +} + export function rawInputToArguments(rawInput: unknown): Record { return isRecord(rawInput) ? rawInput : {}; } diff --git a/ui/desktop/src/acp/chatSessionController.ts b/ui/desktop/src/acp/chatSessionController.ts index ecd4c8c88..af8a88bc9 100644 --- a/ui/desktop/src/acp/chatSessionController.ts +++ b/ui/desktop/src/acp/chatSessionController.ts @@ -81,6 +81,20 @@ function createAcpCreditsExhaustedMessage(error: AcpCreditsExhaustedError): Mess }; } +function assertNoPendingPromptCancellation(sessionId: string): void { + const snapshot = acpChatSessionStore.getSnapshot(sessionId); + if (snapshot?.pendingCancelPromptAttemptId) { + throw new Error('Cannot submit while prompt cancellation is pending'); + } +} + +function assertNoActivePromptAttempt(sessionId: string): void { + const snapshot = acpChatSessionStore.getSnapshot(sessionId); + if (snapshot?.activePromptAttemptId) { + throw new Error('Cannot update message while prompt is active'); + } +} + async function createSession( cwd: string, gooseExtensions: GooseExtension[], @@ -131,7 +145,10 @@ async function submitMessage( userMessage: Message, options: AcpSubmitMessageOptions ): Promise { - if (acpChatSessionStore.getSnapshot(sessionId)?.activePromptAttemptId) { + assertNoPendingPromptCancellation(sessionId); + + const snapshot = acpChatSessionStore.getSnapshot(sessionId); + if (snapshot?.activePromptAttemptId) { return; } @@ -140,10 +157,17 @@ async function submitMessage( try { await acpPromptSession(sessionId, userMessage); + if (acpChatSessionActions.clearPromptCancellation(sessionId, promptAttemptId)) { + return; + } if (acpChatSessionActions.finishPromptAttemptIfCurrent(sessionId, promptAttemptId)) { void options.onFinish(); } } catch (error) { + if (acpChatSessionActions.clearPromptCancellation(sessionId, promptAttemptId)) { + return; + } + const creditsExhaustedError = parseAcpCreditsExhaustedError(error); if (creditsExhaustedError) { if (!acpChatSessionActions.isCurrentPromptAttempt(sessionId, promptAttemptId)) { @@ -175,7 +199,7 @@ function stop(sessionId: string): void { const hasStoredAcpPrompt = storedPromptAttemptId !== null && storedPromptAttemptId !== undefined; if (hasStoredAcpPrompt) { - acpChatSessionActions.clearActivePromptAttempt(sessionId); + acpChatSessionActions.startPromptCancellation(sessionId, storedPromptAttemptId); cancelAcpPermissionRequestsForSession(sessionId); cancelAcpElicitationRequestsForSession(sessionId); acpCancelPrompt(sessionId).catch((error) => { @@ -194,6 +218,9 @@ async function updateMessage( editType: 'fork' | 'edit' | undefined, options: AcpSubmitMessageOptions ): Promise { + assertNoPendingPromptCancellation(sessionId); + assertNoActivePromptAttempt(sessionId); + const resolvedEditType = editType ?? 'fork'; const currentSnapshot = options.getCurrentSnapshot(); diff --git a/ui/desktop/src/acp/chatSessionStore.ts b/ui/desktop/src/acp/chatSessionStore.ts index 7dfc62608..08047dc45 100644 --- a/ui/desktop/src/acp/chatSessionStore.ts +++ b/ui/desktop/src/acp/chatSessionStore.ts @@ -21,12 +21,15 @@ export interface AcpChatSessionSnapshot { chatState: ChatState; sessionLoadError: string | undefined; activePromptAttemptId: string | null; + activeRunId: string | null; + pendingCancelPromptAttemptId: string | null; } type SnapshotListener = (snapshot: AcpChatSessionSnapshot) => void; interface StoreEntry extends AcpChatSessionSnapshot { adapter: AcpSessionNotificationAdapter; + pendingLocalSteerMessageIds: Set; } const initialTokenState: TokenState = { @@ -67,9 +70,18 @@ export interface AcpChatSessionActions { ): AcpChatSessionSnapshot; setMessages(sessionId: string, messages: Message[]): AcpChatSessionSnapshot; + addPendingLocalSteerMessage(sessionId: string, message: Message): AcpChatSessionSnapshot; setChatState(sessionId: string, chatState: ChatState): AcpChatSessionSnapshot; startPromptAttempt(sessionId: string, promptAttemptId: string): AcpChatSessionSnapshot; + startPromptCancellation( + sessionId: string, + promptAttemptId: string + ): AcpChatSessionSnapshot | undefined; + clearPromptCancellation( + sessionId: string, + promptAttemptId: string + ): AcpChatSessionSnapshot | undefined; finishPromptAttemptIfCurrent(sessionId: string, promptAttemptId: string, error?: string): boolean; clearActivePromptAttempt(sessionId: string): AcpChatSessionSnapshot | undefined; isCurrentPromptAttempt(sessionId: string, promptAttemptId: string): boolean; @@ -130,6 +142,9 @@ function createAcpChatSessionStoreInternal(): AcpChatSessionStoreInternal { chatState: ChatState.Idle, sessionLoadError: undefined, activePromptAttemptId: null, + activeRunId: null, + pendingCancelPromptAttemptId: null, + pendingLocalSteerMessageIds: new Set(), adapter: createAcpSessionNotificationAdapter(), }; sessionsById.set(sessionId, entry); @@ -182,7 +197,23 @@ function createAcpChatSessionStoreInternal(): AcpChatSessionStoreInternal { const setMessages: AcpChatSessionActions['setMessages'] = (sessionId, messages) => { const entry = getOrCreateEntry(sessionId); entry.messages = cloneMessages(messages); - entry.adapter = createAcpSessionNotificationAdapter(entry.messages); + retainPendingLocalSteerMessageIds(entry); + entry.adapter = createAdapterForEntry(entry); + return notify(sessionId, entry); + }; + + const addPendingLocalSteerMessage: AcpChatSessionActions['addPendingLocalSteerMessage'] = ( + sessionId, + message + ) => { + const entry = getOrCreateEntry(sessionId); + if (!message.id || entry.messages.some((existing) => existing.id === message.id)) { + return notify(sessionId, entry); + } + + entry.messages = [...entry.messages, cloneMessage(message)]; + entry.pendingLocalSteerMessageIds.add(message.id); + entry.adapter = createAdapterForEntry(entry); return notify(sessionId, entry); }; @@ -206,13 +237,46 @@ function createAcpChatSessionStoreInternal(): AcpChatSessionStoreInternal { promptAttemptId ) => { const entry = getOrCreateEntry(sessionId); + discardPendingLocalSteerMessages(entry); entry.activePromptAttemptId = promptAttemptId; + entry.activeRunId = null; + entry.pendingCancelPromptAttemptId = null; entry.chatState = ChatState.Streaming; entry.sessionLoadError = undefined; entry.notifications = []; return notify(sessionId, entry); }; + const startPromptCancellation: AcpChatSessionActions['startPromptCancellation'] = ( + sessionId, + promptAttemptId + ) => { + const entry = sessionsById.get(sessionId); + if (!entry || entry.activePromptAttemptId !== promptAttemptId) { + return undefined; + } + + entry.activePromptAttemptId = null; + entry.activeRunId = null; + entry.pendingCancelPromptAttemptId = promptAttemptId; + discardPendingLocalSteerMessages(entry); + entry.chatState = ChatState.Idle; + return notify(sessionId, entry); + }; + + const clearPromptCancellation: AcpChatSessionActions['clearPromptCancellation'] = ( + sessionId, + promptAttemptId + ) => { + const entry = sessionsById.get(sessionId); + if (!entry || entry.pendingCancelPromptAttemptId !== promptAttemptId) { + return undefined; + } + + entry.pendingCancelPromptAttemptId = null; + return notify(sessionId, entry); + }; + const finishPromptAttemptIfCurrent: AcpChatSessionActions['finishPromptAttemptIfCurrent'] = ( sessionId, promptAttemptId, @@ -224,6 +288,9 @@ function createAcpChatSessionStoreInternal(): AcpChatSessionStoreInternal { } entry.activePromptAttemptId = null; + entry.activeRunId = null; + entry.pendingCancelPromptAttemptId = null; + discardPendingLocalSteerMessages(entry); entry.chatState = ChatState.Idle; entry.sessionLoadError = error; notify(sessionId, entry); @@ -239,6 +306,8 @@ function createAcpChatSessionStoreInternal(): AcpChatSessionStoreInternal { } entry.activePromptAttemptId = null; + entry.activeRunId = null; + discardPendingLocalSteerMessages(entry); entry.chatState = ChatState.Idle; return notify(sessionId, entry); }; @@ -310,8 +379,11 @@ function createAcpChatSessionStoreInternal(): AcpChatSessionStoreInternal { failSessionLoad, setSessionLoadError, setMessages, + addPendingLocalSteerMessage, setChatState, startPromptAttempt, + startPromptCancellation, + clearPromptCancellation, finishPromptAttemptIfCurrent, clearActivePromptAttempt, isCurrentPromptAttempt, @@ -382,8 +454,11 @@ function actionsFromStore(store: AcpChatSessionStoreInternal): AcpChatSessionAct failSessionLoad: store.failSessionLoad, setSessionLoadError: store.setSessionLoadError, setMessages: store.setMessages, + addPendingLocalSteerMessage: store.addPendingLocalSteerMessage, setChatState: store.setChatState, startPromptAttempt: store.startPromptAttempt, + startPromptCancellation: store.startPromptCancellation, + clearPromptCancellation: store.clearPromptCancellation, finishPromptAttemptIfCurrent: store.finishPromptAttemptIfCurrent, clearActivePromptAttempt: store.clearActivePromptAttempt, isCurrentPromptAttempt: store.isCurrentPromptAttempt, @@ -395,6 +470,7 @@ function applyChatStateChanges(entry: StoreEntry, changes: AcpChatStateChange[]) switch (change.type) { case 'messages': entry.messages = cloneMessages(change.messages); + retainPendingLocalSteerMessageIds(entry); break; case 'tokenState': entry.tokenState = { ...entry.tokenState, ...change.tokenState }; @@ -403,6 +479,12 @@ function applyChatStateChanges(entry: StoreEntry, changes: AcpChatStateChange[]) if (change.name && entry.session) { entry.session = { ...entry.session, name: change.name }; } + if (change.activeRunId !== undefined) { + entry.activeRunId = change.activeRunId; + } + break; + case 'localSteerConfirmed': + entry.pendingLocalSteerMessageIds.delete(change.messageId); break; case 'notification': entry.notifications = [...entry.notifications, change.notification]; @@ -415,9 +497,63 @@ function resetReplayState(entry: StoreEntry): void { entry.messages = []; entry.tokenState = { ...initialTokenState }; entry.notifications = []; + entry.activeRunId = null; + entry.pendingCancelPromptAttemptId = null; + entry.pendingLocalSteerMessageIds.clear(); entry.adapter = createAcpSessionNotificationAdapter(); } +function retainPendingLocalSteerMessageIds(entry: StoreEntry): void { + if (entry.pendingLocalSteerMessageIds.size === 0) { + return; + } + + const messageIds = new Set(entry.messages.map((message) => message.id).filter(Boolean)); + entry.pendingLocalSteerMessageIds = new Set( + [...entry.pendingLocalSteerMessageIds].filter((messageId) => messageIds.has(messageId)) + ); +} + +function discardPendingLocalSteerMessages(entry: StoreEntry): void { + if (entry.pendingLocalSteerMessageIds.size === 0) { + return; + } + + entry.messages = entry.messages.filter( + (message) => !message.id || !entry.pendingLocalSteerMessageIds.has(message.id) + ); + entry.pendingLocalSteerMessageIds.clear(); + entry.adapter = createAdapterForEntry(entry); +} + +function createAdapterForEntry(entry: StoreEntry): AcpSessionNotificationAdapter { + return createAcpSessionNotificationAdapter( + entry.messages, + confirmedLocalSteerTextByMessageId(entry) + ); +} + +function confirmedLocalSteerTextByMessageId(entry: StoreEntry): Map { + const textByMessageId = new Map(); + + for (const message of entry.messages) { + if ( + !message.id || + !message.metadata.steer || + entry.pendingLocalSteerMessageIds.has(message.id) + ) { + continue; + } + + const firstContent = message.content[0]; + if (firstContent?.type === 'text') { + textByMessageId.set(message.id, firstContent.text); + } + } + + return textByMessageId; +} + function snapshotFromEntry(entry: StoreEntry): AcpChatSessionSnapshot { return { session: entry.session, @@ -427,6 +563,8 @@ function snapshotFromEntry(entry: StoreEntry): AcpChatSessionSnapshot { chatState: entry.chatState, sessionLoadError: entry.sessionLoadError, activePromptAttemptId: entry.activePromptAttemptId, + activeRunId: entry.activeRunId, + pendingCancelPromptAttemptId: entry.pendingCancelPromptAttemptId, }; } diff --git a/ui/desktop/src/acp/prompt.ts b/ui/desktop/src/acp/prompt.ts index 780fba242..5c84af2f8 100644 --- a/ui/desktop/src/acp/prompt.ts +++ b/ui/desktop/src/acp/prompt.ts @@ -1,4 +1,5 @@ import type { ContentBlock, PromptResponse } from '@agentclientprotocol/sdk'; +import type { SteerSessionRequest_unstable, SteerSessionResponse_unstable } from '@aaif/goose-sdk'; import type { Message } from '../api'; import { getAcpClient } from './acpConnection'; @@ -18,6 +19,19 @@ export async function acpCancelPrompt(sessionId: string): Promise { await client.cancel({ sessionId }); } +export async function acpSteerSession( + sessionId: string, + message: Message, + expectedRunId: string +): Promise { + const client = await getAcpClient(); + return client.goose.sessionSteer_unstable({ + sessionId, + expectedRunId, + prompt: messageToAcpPromptContent(message) as unknown as SteerSessionRequest_unstable['prompt'], + }); +} + export function messageToAcpPromptContent(message: Message): ContentBlock[] { const prompt: ContentBlock[] = []; diff --git a/ui/desktop/src/acp/sessionNotificationAdapter.ts b/ui/desktop/src/acp/sessionNotificationAdapter.ts index 3548a5084..27d7d566a 100644 --- a/ui/desktop/src/acp/sessionNotificationAdapter.ts +++ b/ui/desktop/src/acp/sessionNotificationAdapter.ts @@ -9,7 +9,12 @@ import { import { applyGooseSessionNotification } from './adapter/gooseSessionNotifications'; import { applyContentChunk, applyThoughtChunk } from './adapter/messages'; import { applyPermissionRequest as applyPermissionRequestToState } from './adapter/permissions'; -import { type AcpChatStateChange, type AdapterState, cloneMessage } from './adapter/shared'; +import { + type AcpChatStateChange, + type AdapterState, + cloneMessage, + getGooseActiveRunId, +} from './adapter/shared'; import { applyToolCall, applyToolCallUpdate } from './adapter/tools'; import type { AcpElicitationRequest } from './elicitationRequests'; @@ -25,10 +30,12 @@ export interface AcpSessionNotificationAdapter { } export function createAcpSessionNotificationAdapter( - initialMessages: Message[] = [] + initialMessages: Message[] = [], + localSteerTextByMessageId: Map = new Map() ): AcpSessionNotificationAdapter { const state: AdapterState = { messages: initialMessages.map(cloneMessage), + localSteerTextByMessageId: new Map(localSteerTextByMessageId), }; return { @@ -70,13 +77,20 @@ function applyAcpSessionNotification( return applyToolCall(state, update); case 'tool_call_update': return applyToolCallUpdate(state, update); - case 'session_info_update': + case 'session_info_update': { + const activeRunId = getGooseActiveRunId(update); + if (!update.title && activeRunId === undefined) { + return []; + } + return [ { type: 'sessionInfo', ...(update.title ? { name: update.title } : {}), + ...(activeRunId !== undefined ? { activeRunId } : {}), }, ]; + } case 'usage_update': return []; default: diff --git a/ui/desktop/src/acpChatFeatureFlag.ts b/ui/desktop/src/acpChatFeatureFlag.ts index edea9d78a..34ebb1753 100644 --- a/ui/desktop/src/acpChatFeatureFlag.ts +++ b/ui/desktop/src/acpChatFeatureFlag.ts @@ -1 +1 @@ -export const USE_ACP_CHAT = false; \ No newline at end of file +export const USE_ACP_CHAT = false; diff --git a/ui/desktop/src/components/BaseChat.tsx b/ui/desktop/src/components/BaseChat.tsx index 9519945bb..28b3adc2e 100644 --- a/ui/desktop/src/components/BaseChat.tsx +++ b/ui/desktop/src/components/BaseChat.tsx @@ -95,6 +95,7 @@ export default function BaseChat({ setChatState, updateSession, handleSubmit, + onSteerQueuedMessage, submitElicitationResponse, stopStreaming, sessionLoadError, @@ -102,6 +103,7 @@ export default function BaseChat({ tokenState, notifications: toolCallNotifications, pauseQueueOnStop, + queueProcessingBlocked, onMessageUpdate, } = useChatSession({ sessionId, @@ -505,7 +507,9 @@ export default function BaseChat({ chatState={chatState} setChatState={setChatState} onStop={stopStreaming} + onSteerQueuedMessage={onSteerQueuedMessage} pauseQueueOnStop={pauseQueueOnStop} + queueProcessingBlocked={queueProcessingBlocked} commandHistory={commandHistory} initialValue={initialPrompt} setView={setView} diff --git a/ui/desktop/src/components/ChatInput.tsx b/ui/desktop/src/components/ChatInput.tsx index d1f7e264c..75e1b2a36 100644 --- a/ui/desktop/src/components/ChatInput.tsx +++ b/ui/desktop/src/components/ChatInput.tsx @@ -154,6 +154,10 @@ const i18n = defineMessages({ id: 'chatInput.send', defaultMessage: 'Send', }, + waitingForCancellation: { + id: 'chatInput.waitingForCancellation', + defaultMessage: 'Waiting for cancellation to finish', + }, failedToReadImage: { id: 'chatInput.failedToReadImage', defaultMessage: 'Failed to read image file', @@ -170,7 +174,9 @@ interface ChatInputProps { chatState: ChatState; setChatState?: (state: ChatState) => void; onStop?: () => void; + onSteerQueuedMessage?: (input: UserInput) => Promise; pauseQueueOnStop?: boolean; + queueProcessingBlocked?: boolean; commandHistory?: string[]; initialValue?: string; droppedFiles?: DroppedFile[]; @@ -205,7 +211,9 @@ export default function ChatInput({ chatState = ChatState.Idle, setChatState, onStop, + onSteerQueuedMessage, pauseQueueOnStop = false, + queueProcessingBlocked = false, commandHistory = [], initialValue = '', droppedFiles = [], @@ -241,15 +249,35 @@ export default function ChatInput({ // Derived state - chatState != Idle means we're in some form of loading state const isLoading = chatState !== ChatState.Idle; + const isLoadingRef = useRef(isLoading); + const queueProcessingBlockedRef = useRef(queueProcessingBlocked); const wasLoadingRef = useRef(isLoading); + const wasQueueProcessingBlockedRef = useRef(queueProcessingBlocked); + isLoadingRef.current = isLoading; + queueProcessingBlockedRef.current = queueProcessingBlocked; // Queue functionality - ephemeral, only exists in memory for this chat instance const [queuedMessages, setQueuedMessages] = useState([]); const queuePausedRef = useRef(false); const editingMessageIdRef = useRef(null); const sendAfterStopMessageIdRef = useRef(null); + const sendNowInFlightMessageIdsRef = useRef>(new Set()); + const [sendNowInFlightMessageIds, setSendNowInFlightMessageIds] = useState>( + new Set() + ); const [lastInterruption, setLastInterruption] = useState(null); + const setSendNowInFlightMessage = useCallback((messageId: string, isInFlight: boolean) => { + const nextMessageIds = new Set(sendNowInFlightMessageIdsRef.current); + if (isInFlight) { + nextMessageIds.add(messageId); + } else { + nextMessageIds.delete(messageId); + } + sendNowInFlightMessageIdsRef.current = nextMessageIds; + setSendNowInFlightMessageIds(nextMessageIds); + }, []); + const pauseRemainingQueue = useCallback(() => { queuePausedRef.current = true; }, []); @@ -357,7 +385,16 @@ export default function ChatInput({ // Queue processing useEffect(() => { - if (wasLoadingRef.current && !isLoading && queuedMessages.length > 0) { + const becameIdle = wasLoadingRef.current && !isLoading; + const becameUnblocked = wasQueueProcessingBlockedRef.current && !queueProcessingBlocked; + const hasSendNowInFlight = sendNowInFlightMessageIdsRef.current.size > 0; + + if ( + (becameIdle || (becameUnblocked && !isLoading)) && + !queueProcessingBlocked && + !hasSendNowInFlight && + queuedMessages.length > 0 + ) { const pendingSendAfterStopId = sendAfterStopMessageIdRef.current; const messageToSend = pendingSendAfterStopId ? queuedMessages.find((message) => message.id === pendingSendAfterStopId) @@ -366,11 +403,13 @@ export default function ChatInput({ if (pendingSendAfterStopId && !messageToSend) { clearPendingSendAfterStop(pendingSendAfterStopId); wasLoadingRef.current = isLoading; + wasQueueProcessingBlockedRef.current = queueProcessingBlocked; return; } if (!messageToSend) { wasLoadingRef.current = isLoading; + wasQueueProcessingBlockedRef.current = queueProcessingBlocked; return; } @@ -406,8 +445,10 @@ export default function ChatInput({ } } wasLoadingRef.current = isLoading; + wasQueueProcessingBlockedRef.current = queueProcessingBlocked; }, [ isLoading, + queueProcessingBlocked, queuedMessages, handleSubmit, lastInterruption, @@ -1074,6 +1115,7 @@ export default function ChatInput({ const canSubmit = !isLoading && + !queueProcessingBlocked && (displayValue.trim() || pastedImages.some((img) => img.dataUrl && !img.error && !img.isLoading) || allDroppedFiles.some((file) => !file.error && !file.isLoading)); @@ -1190,12 +1232,16 @@ export default function ChatInput({ const onFormSubmit = (e: React.FormEvent | React.MouseEvent) => { e.preventDefault(); + if (queueProcessingBlocked) { + return; + } if (isLoading && hasSubmittableContent) { handleInterruptionAndQueue(); return; } const canSubmit = !isLoading && + !queueProcessingBlocked && (displayValue.trim() || pastedImages.some((img) => img.dataUrl && !img.error && !img.isLoading) || allDroppedFiles.some((file) => !file.error && !file.isLoading)); @@ -1314,9 +1360,11 @@ export default function ChatInput({ isAnyDroppedFileLoading || isRecording || isTranscribing || + queueProcessingBlocked || chatState === ChatState.RestartingAgent; const getSubmitButtonTooltip = (): string => { + if (queueProcessingBlocked) return intl.formatMessage(i18n.waitingForCancellation); if (isAnyImageLoading) return intl.formatMessage(i18n.waitingForImages); if (isAnyDroppedFileLoading) return intl.formatMessage(i18n.processingDroppedFiles); if (isRecording) return intl.formatMessage(i18n.recording); @@ -1328,28 +1376,35 @@ export default function ChatInput({ // Queue management functions - no storage persistence, only in-memory const handleRemoveQueuedMessage = (messageId: string) => { + if (sendNowInFlightMessageIdsRef.current.has(messageId)) return; clearPendingSendAfterStop(messageId); setQueuedMessages((prev) => prev.filter((msg) => msg.id !== messageId)); }; const handleClearQueue = () => { + if (sendNowInFlightMessageIdsRef.current.size > 0) return; setQueuedMessages([]); clearQueueState(); }; const handleReorderMessages = (reorderedMessages: QueuedMessage[]) => { + if (reorderedMessages.some((message) => sendNowInFlightMessageIdsRef.current.has(message.id))) { + return; + } setQueuedMessages(reorderedMessages); }; const handleEditMessage = (messageId: string, newContent: string) => { + if (sendNowInFlightMessageIdsRef.current.has(messageId)) return; setQueuedMessages((prev) => prev.map((msg) => (msg.id === messageId ? { ...msg, content: newContent } : msg)) ); }; - const handleStopAndSend = (messageId: string) => { + const handleStopAndSend = async (messageId: string) => { const messageToSend = queuedMessages.find((msg) => msg.id === messageId); if (!messageToSend) return; + if (queueProcessingBlocked) return; if (!isLoading) { setQueuedMessages((prev) => removeQueuedMessage(prev, messageId)); @@ -1358,6 +1413,53 @@ export default function ChatInput({ return; } + if (onSteerQueuedMessage) { + if (sendNowInFlightMessageIdsRef.current.has(messageId)) { + return; + } + + const wasQueuePausedBeforeSteer = queuePausedRef.current; + pauseRemainingQueue(); + setSendNowInFlightMessage(messageId, true); + try { + const steerAccepted = await onSteerQueuedMessage({ + msg: messageToSend.content, + images: messageToSend.images, + }); + + if (steerAccepted) { + LocalMessageStorage.addMessage(messageToSend.content); + clearPendingSendAfterStop(messageId); + setQueuedMessages((prev) => { + const newQueue = removeQueuedMessage(prev, messageId); + if (newQueue.length === 0) { + clearQueueState(); + } else { + pauseRemainingQueue(); + } + return newQueue; + }); + return; + } + } finally { + setSendNowInFlightMessage(messageId, false); + } + + if (!isLoadingRef.current && !queueProcessingBlockedRef.current) { + queuePausedRef.current = wasQueuePausedBeforeSteer; + setQueuedMessages((prev) => { + const newQueue = removeQueuedMessage(prev, messageId); + if (newQueue.length === 0) { + clearQueueState(); + } + return newQueue; + }); + LocalMessageStorage.addMessage(messageToSend.content); + handleSubmit({ msg: messageToSend.content, images: messageToSend.images }); + return; + } + } + sendAfterStopMessageIdRef.current = messageId; pauseRemainingQueue(); setQueuedMessages((prev) => moveQueuedMessageToFront(prev, messageId)); @@ -1374,7 +1476,7 @@ export default function ChatInput({ const handleResumeQueue = () => { queuePausedRef.current = false; setLastInterruption(null); - if (!isLoading && queuedMessages.length > 0) { + if (!isLoading && !queueProcessingBlocked && queuedMessages.length > 0) { const nextMessage = queuedMessages[0]; LocalMessageStorage.addMessage(nextMessage.content); handleSubmit({ msg: nextMessage.content, images: nextMessage.images }); @@ -1421,6 +1523,7 @@ export default function ChatInput({ onEditMessage={handleEditMessage} onTriggerQueueProcessing={handleResumeQueue} editingMessageIdRef={editingMessageIdRef} + sendingMessageIds={sendNowInFlightMessageIds} isPaused={queuePausedRef.current} className="border-b border-border-primary" /> diff --git a/ui/desktop/src/components/MessageQueue.tsx b/ui/desktop/src/components/MessageQueue.tsx index ce856ae74..242e038e8 100644 --- a/ui/desktop/src/components/MessageQueue.tsx +++ b/ui/desktop/src/components/MessageQueue.tsx @@ -103,6 +103,7 @@ interface MessageQueueProps { onTriggerQueueProcessing?: () => void; editingMessageIdRef?: React.MutableRefObject; onReorderMessages?: (reorderedMessages: QueuedMessage[]) => void; + sendingMessageIds?: ReadonlySet; className?: string; isPaused?: boolean; } @@ -116,6 +117,7 @@ export const MessageQueue: React.FC = ({ onTriggerQueueProcessing, editingMessageIdRef, onReorderMessages, + sendingMessageIds, className = '', isPaused = false, }) => { @@ -126,18 +128,28 @@ export const MessageQueue: React.FC = ({ const [hoveredMessage, setHoveredMessage] = useState(null); const [editingMessage, setEditingMessage] = useState(null); const [editContent, setEditContent] = useState(''); + const isSendingMessage = (messageId: string) => sendingMessageIds?.has(messageId) ?? false; if (queuedMessages.length === 0) { return null; } const handleDragStart = (e: React.DragEvent, messageId: string) => { + if (isSendingMessage(messageId)) { + e.preventDefault(); + return; + } + setDraggedItem(messageId); e.dataTransfer.effectAllowed = 'move'; e.dataTransfer.setData('text/html', messageId); }; const handleDragOver = (e: React.DragEvent, messageId: string) => { + if (isSendingMessage(messageId)) { + return; + } + e.preventDefault(); e.dataTransfer.dropEffect = 'move'; setDragOverItem(messageId); @@ -150,7 +162,7 @@ export const MessageQueue: React.FC = ({ const handleDrop = (e: React.DragEvent, targetMessageId: string) => { e.preventDefault(); - if (!draggedItem || !onReorderMessages) return; + if (!draggedItem || !onReorderMessages || isSendingMessage(targetMessageId)) return; const draggedIndex = queuedMessages.findIndex((msg) => msg.id === draggedItem); const targetIndex = queuedMessages.findIndex((msg) => msg.id === targetMessageId); @@ -185,6 +197,8 @@ export const MessageQueue: React.FC = ({ const nextMessage = queuedMessages[0]; const remainingCount = queuedMessages.length - 1; + const nextMessageIsSending = isSendingMessage(nextMessage.id); + const hasSendingMessages = queuedMessages.some((message) => isSendingMessage(message.id)); // Compact View if (!isExpanded) { @@ -232,8 +246,10 @@ export const MessageQueue: React.FC = ({ size="sm" onClick={(e) => { e.stopPropagation(); + if (nextMessageIsSending) return; onStopAndSend(nextMessage.id); }} + disabled={nextMessageIsSending} className="h-7 px-2 text-xs text-info hover:text-info/80 hover:bg-info/10" title={intl.formatMessage(i18n.sendNow)} > @@ -287,12 +303,16 @@ export const MessageQueue: React.FC = ({
- {isPaused ? intl.formatMessage(i18n.queuePaused) : intl.formatMessage(i18n.messageQueue)} + {isPaused + ? intl.formatMessage(i18n.queuePaused) + : intl.formatMessage(i18n.messageQueue)} {intl.formatMessage(i18n.messageCount, { count: queuedMessages.length, - status: isPaused ? intl.formatMessage(i18n.waiting) : intl.formatMessage(i18n.queued), + status: isPaused + ? intl.formatMessage(i18n.waiting) + : intl.formatMessage(i18n.queued), })}
@@ -304,6 +324,7 @@ export const MessageQueue: React.FC = ({ variant="ghost" size="sm" onClick={onClearQueue} + disabled={hasSendingMessages} className="text-xs h-7 px-3 text-muted-foreground hover:text-destructive hover:bg-destructive/10 transition-colors" > {intl.formatMessage(i18n.clearAll)} @@ -328,184 +349,193 @@ export const MessageQueue: React.FC = ({
- - {intl.formatMessage(i18n.queuePausedExpanded)} - + {intl.formatMessage(i18n.queuePausedExpanded)}
)} {/* Message Bubbles */}
- {queuedMessages.map((message, index) => ( -
handleDragStart(e, message.id)} - onDragOver={(e) => handleDragOver(e, message.id)} - onDragLeave={handleDragLeave} - onDrop={(e) => handleDrop(e, message.id)} - onDragEnd={handleDragEnd} - onMouseEnter={() => setHoveredMessage(message.id)} - onMouseLeave={() => setHoveredMessage(null)} - > - {/* Main message bubble */} + {queuedMessages.map((message, index) => { + const isSending = isSendingMessage(message.id); + const isEditing = editingMessage === message.id; + return (
handleDragStart(e, message.id)} + onDragOver={(e) => handleDragOver(e, message.id)} + onDragLeave={handleDragLeave} + onDrop={(e) => handleDrop(e, message.id)} + onDragEnd={handleDragEnd} + onMouseEnter={() => setHoveredMessage(message.id)} + onMouseLeave={() => setHoveredMessage(null)} > - {/* Priority indicator */} -
-
- {index + 1} -
- - {/* Drag handle */} - {onReorderMessages && ( + {/* Main message bubble */} +
+ {/* Priority indicator */} +
- + {index + 1}
- )} -
- {/* Message content */} -
- {editingMessage === message.id ? ( -
-