From ac53ec0bdcf21516b56610f3a4222b88203cdcf0 Mon Sep 17 00:00:00 2001 From: Bohdan Triapitsyn Date: Sun, 21 Dec 2025 00:47:43 +0200 Subject: [PATCH] feat(chat): add revert to previous message Add revert button on user messages except first message Restore reverted message text to input field Filter out messages after revert point --- packages/ui/src/components/chat/ChatInput.tsx | 16 ++++ .../ui/src/components/chat/ChatMessage.tsx | 11 +++ .../ui/src/components/chat/MessageList.tsx | 33 ++++--- .../components/chat/message/MessageBody.tsx | 91 ++++++++++++------- packages/ui/src/lib/opencode/client.ts | 19 ++++ packages/ui/src/stores/messageStore.ts | 41 ++++++++- packages/ui/src/stores/types/sessionTypes.ts | 6 ++ packages/ui/src/stores/useSessionStore.ts | 58 ++++++++++++ 8 files changed, 227 insertions(+), 48 deletions(-) diff --git a/packages/ui/src/components/chat/ChatInput.tsx b/packages/ui/src/components/chat/ChatInput.tsx index a0b38d3c..27b24a6f 100644 --- a/packages/ui/src/components/chat/ChatInput.tsx +++ b/packages/ui/src/components/chat/ChatInput.tsx @@ -71,6 +71,8 @@ export const ChatInput: React.FC = ({ onOpenSettings, scrollToBo const addServerFile = useSessionStore((state) => state.addServerFile); const clearAttachedFiles = useSessionStore((state) => state.clearAttachedFiles); const saveSessionAgentSelection = useSessionStore((state) => state.saveSessionAgentSelection); + const consumePendingInputText = useSessionStore((state) => state.consumePendingInputText); + const pendingInputText = useSessionStore((state) => state.pendingInputText); const { currentProviderId, currentModelId, currentAgentName, setAgent, getVisibleAgents } = useConfigStore(); const agents = getVisibleAgents(); @@ -125,6 +127,20 @@ export const ChatInput: React.FC = ({ onOpenSettings, scrollToBo } }, [isMobile]); + // Consume pending input text (e.g., from revert action) + React.useEffect(() => { + if (pendingInputText !== null) { + const text = consumePendingInputText(); + if (text) { + setMessage(text); + // Focus textarea after setting message + setTimeout(() => { + textareaRef.current?.focus(); + }, 0); + } + } + }, [pendingInputText, consumePendingInputText]); + const currentAgent = React.useMemo(() => { if (!currentAgentName) { return undefined; diff --git a/packages/ui/src/components/chat/ChatMessage.tsx b/packages/ui/src/components/chat/ChatMessage.tsx index 4fb07dda..77fe5811 100644 --- a/packages/ui/src/components/chat/ChatMessage.tsx +++ b/packages/ui/src/components/chat/ChatMessage.tsx @@ -66,6 +66,7 @@ interface ChatMessageProps { scrollToBottom?: (options?: { instant?: boolean; force?: boolean }) => void; isPendingAnchor?: boolean; turnGroupingContext?: TurnGroupingContext; + isFirstMessage?: boolean; } const ChatMessage: React.FC = ({ @@ -76,6 +77,7 @@ const ChatMessage: React.FC = ({ animationHandlers, isPendingAnchor = false, turnGroupingContext, + isFirstMessage = false, }) => { const { isMobile, hasTouchInput } = useDeviceInfo(); const { currentTheme } = useThemeSystem(); @@ -529,6 +531,13 @@ const ChatMessage: React.FC = ({ setTimeout(() => setCopiedMessage(false), 2000); }, [messageTextContent]); + const revertToMessage = useSessionStore((state) => state.revertToMessage); + + const handleRevert = React.useCallback(() => { + if (!sessionId || !message.info.id) return; + revertToMessage(sessionId, message.info.id); + }, [sessionId, message.info.id, revertToMessage]); + const handleToggleTool = React.useCallback((toolId: string) => { setExpandedTools((prev) => { const next = new Set(prev); @@ -731,6 +740,8 @@ const ChatMessage: React.FC = ({ showReasoningTraces={showReasoningTraces} onAuxiliaryContentComplete={handleAuxiliaryContentComplete} agentMention={agentMention} + onRevert={handleRevert} + isFirstMessage={isFirstMessage} /> diff --git a/packages/ui/src/components/chat/MessageList.tsx b/packages/ui/src/components/chat/MessageList.tsx index 567f6a35..33ba9244 100644 --- a/packages/ui/src/components/chat/MessageList.tsx +++ b/packages/ui/src/components/chat/MessageList.tsx @@ -75,19 +75,26 @@ const MessageList: React.FC = ({ )}
- {displayMessages.map((message, index) => ( - 0 ? displayMessages[index - 1] : undefined} - nextMessage={index < displayMessages.length - 1 ? displayMessages[index + 1] : undefined} - onContentChange={onMessageContentChange} - animationHandlers={getAnimationHandlers(message.info.id)} - scrollToBottom={scrollToBottom} - isPendingAnchor={pendingAnchorId === message.info.id} - turnGroupingContext={getContextForMessage(message.info.id)} - /> - ))} + {displayMessages.map((message, index) => { + // Check if this is the first user message + const isFirstUserMessage = message.info.role === 'user' && + !displayMessages.slice(0, index).some((m) => m.info.role === 'user'); + + return ( + 0 ? displayMessages[index - 1] : undefined} + nextMessage={index < displayMessages.length - 1 ? displayMessages[index + 1] : undefined} + onContentChange={onMessageContentChange} + animationHandlers={getAnimationHandlers(message.info.id)} + scrollToBottom={scrollToBottom} + isPendingAnchor={pendingAnchorId === message.info.id} + turnGroupingContext={getContextForMessage(message.info.id)} + isFirstMessage={isFirstUserMessage} + /> + ); + })}
diff --git a/packages/ui/src/components/chat/message/MessageBody.tsx b/packages/ui/src/components/chat/message/MessageBody.tsx index ffe86b37..3c86808b 100644 --- a/packages/ui/src/components/chat/message/MessageBody.tsx +++ b/packages/ui/src/components/chat/message/MessageBody.tsx @@ -16,7 +16,7 @@ import { isEmptyTextPart, extractTextContent } from './partUtils'; import { FadeInOnReveal } from './FadeInOnReveal'; import { Button } from '@/components/ui/button'; import { Tooltip, TooltipContent, TooltipTrigger } from '@/components/ui/tooltip'; -import { RiCheckLine, RiFileCopyLine, RiChatNewLine } from '@remixicon/react'; +import { RiCheckLine, RiFileCopyLine, RiChatNewLine, RiArrowGoBackLine } from '@remixicon/react'; import type { ContentChangeReason } from '@/hooks/useChatScrollManager'; import { SimpleMarkdownRenderer } from '../MarkdownRenderer'; @@ -127,6 +127,8 @@ interface MessageBodyProps { showReasoningTraces?: boolean; agentMention?: AgentMentionInfo; turnGroupingContext?: TurnGroupingContext; + onRevert?: () => void; + isFirstMessage?: boolean; } const UserMessageBody: React.FC<{ @@ -139,7 +141,9 @@ const UserMessageBody: React.FC<{ copiedMessage?: boolean; onShowPopup: (content: ToolPopupContent) => void; agentMention?: AgentMentionInfo; -}> = ({ messageId, parts, isMobile, hasTouchInput, hasTextContent, onCopyMessage, copiedMessage, onShowPopup, agentMention }) => { + onRevert?: () => void; + isFirstMessage?: boolean; +}> = ({ messageId, parts, isMobile, hasTouchInput, hasTextContent, onCopyMessage, copiedMessage, onShowPopup, agentMention, onRevert, isFirstMessage }) => { const [copyHintVisible, setCopyHintVisible] = React.useState(false); const copyHintTimeoutRef = React.useRef(null); @@ -231,40 +235,63 @@ const UserMessageBody: React.FC<{ })} - {canCopyMessage && hasCopyableText && ( + {(canCopyMessage && hasCopyableText) || (onRevert && !isFirstMessage) ? (
- - - - - Copy message - + {onRevert && !isFirstMessage && ( + + + + + Revert from here + + )} + {canCopyMessage && hasCopyableText && ( + + + + + Copy message + + )}
- )} + ) : null} ); }; @@ -1193,6 +1220,8 @@ const MessageBody: React.FC = ({ isUser, ...props }) => { copiedMessage={props.copiedMessage} onShowPopup={props.onShowPopup} agentMention={props.agentMention} + onRevert={props.onRevert} + isFirstMessage={props.isFirstMessage} /> ); } diff --git a/packages/ui/src/lib/opencode/client.ts b/packages/ui/src/lib/opencode/client.ts index 732403fd..7012336b 100644 --- a/packages/ui/src/lib/opencode/client.ts +++ b/packages/ui/src/lib/opencode/client.ts @@ -478,6 +478,25 @@ class OpencodeService { return Boolean(response.data); } + async revertSession(sessionId: string, messageId: string, partId?: string): Promise { + const response = await this.client.session.revert({ + path: { id: sessionId }, + query: this.currentDirectory ? { directory: this.currentDirectory } : undefined, + body: { messageID: messageId, partID: partId } + }); + if (!response.data) throw new Error('Failed to revert session'); + return response.data; + } + + async unrevertSession(sessionId: string): Promise { + const response = await this.client.session.unrevert({ + path: { id: sessionId }, + query: this.currentDirectory ? { directory: this.currentDirectory } : undefined + }); + if (!response.data) throw new Error('Failed to unrevert session'); + return response.data; + } + async getSessionStatus(): Promise< Record > { diff --git a/packages/ui/src/stores/messageStore.ts b/packages/ui/src/stores/messageStore.ts index dcce78aa..b5231e9a 100644 --- a/packages/ui/src/stores/messageStore.ts +++ b/packages/ui/src/stores/messageStore.ts @@ -290,6 +290,30 @@ const resolveSessionDirectory = async (sessionId: string | null | undefined): Pr } }; +const getSessionRevertMessageId = (sessionId: string | null | undefined): string | null => { + if (!sessionId) return null; + try { + const sessionStore = useSessionStore.getState(); + const session = sessionStore.sessions.find((entry) => entry.id === sessionId) as { revert?: { messageID?: string } } | undefined; + return session?.revert?.messageID ?? null; + } catch { + return null; + } +}; + +const filterRevertedMessages = ( + messages: { info: Message; parts: Part[] }[], + revertMessageId: string | null +): { info: Message; parts: Part[] }[] => { + if (!revertMessageId) return messages; + + const revertIndex = messages.findIndex((m) => m.info.id === revertMessageId); + if (revertIndex === -1) return messages; + + // Keep only messages before the revert point (exclusive) + return messages.slice(0, revertIndex); +}; + const executeWithSessionDirectory = async (sessionId: string | null | undefined, operation: () => Promise): Promise => { const directoryOverride = await resolveSessionDirectory(sessionId); if (directoryOverride) { @@ -358,15 +382,20 @@ export const useMessageStore = create()( loadMessages: async (sessionId: string, limit: number = MEMORY_LIMITS.VIEWPORT_MESSAGES) => { const allMessages = await executeWithSessionDirectory(sessionId, () => opencodeClient.getSessionMessages(sessionId)); + + // Filter out reverted messages first + const revertMessageId = getSessionRevertMessageId(sessionId); + const messagesWithoutReverted = filterRevertedMessages(allMessages, revertMessageId); + const watermark = get().sessionMemoryState.get(sessionId)?.trimmedHeadMaxId; const afterWatermark = watermark - ? allMessages.filter((message) => { + ? messagesWithoutReverted.filter((message) => { const messageId = message?.info?.id; if (!messageId) return true; return isIdNewer(messageId, watermark); }) - : allMessages; + : messagesWithoutReverted; const messagesToKeep = afterWatermark.slice(-limit); set((state) => { @@ -1836,15 +1865,19 @@ export const useMessageStore = create()( }, syncMessages: (sessionId: string, messages: { info: Message; parts: Part[] }[]) => { + // Filter out reverted messages first + const revertMessageId = getSessionRevertMessageId(sessionId); + const messagesWithoutReverted = filterRevertedMessages(messages, revertMessageId); + const watermark = get().sessionMemoryState.get(sessionId)?.trimmedHeadMaxId; const messagesFiltered = watermark - ? messages.filter((message) => { + ? messagesWithoutReverted.filter((message) => { const messageId = message?.info?.id; if (!messageId) return true; return isIdNewer(messageId, watermark); }) - : messages; + : messagesWithoutReverted; set((state) => { const newMessages = new Map(state.messages); diff --git a/packages/ui/src/stores/types/sessionTypes.ts b/packages/ui/src/stores/types/sessionTypes.ts index 37c3f9e5..a7295519 100644 --- a/packages/ui/src/stores/types/sessionTypes.ts +++ b/packages/ui/src/stores/types/sessionTypes.ts @@ -96,6 +96,8 @@ export interface SessionStore { userSummaryTitles: Map; + pendingInputText: string | null; + getSessionAgentEditMode: (sessionId: string, agentName: string | undefined, defaultMode?: EditPermissionMode) => EditPermissionMode; toggleSessionAgentEditMode: (sessionId: string, agentName: string | undefined, defaultMode?: EditPermissionMode) => void; setSessionAgentEditMode: (sessionId: string, agentName: string | undefined, mode: EditPermissionMode, defaultMode?: EditPermissionMode) => void; @@ -169,4 +171,8 @@ export interface SessionStore { pollForTokenUpdates: (sessionId: string, messageId: string, maxAttempts?: number) => void; updateSession: (session: Session) => void; + + revertToMessage: (sessionId: string, messageId: string) => Promise; + setPendingInputText: (text: string | null) => void; + consumePendingInputText: () => string | null; } diff --git a/packages/ui/src/stores/useSessionStore.ts b/packages/ui/src/stores/useSessionStore.ts index 65b959a5..d5bc77fd 100644 --- a/packages/ui/src/stores/useSessionStore.ts +++ b/packages/ui/src/stores/useSessionStore.ts @@ -94,6 +94,7 @@ export const useSessionStore = create()( abortPromptExpiresAt: null, sessionActivityPhase: new Map(), userSummaryTitles: new Map(), + pendingInputText: null, getSessionAgentEditMode: (sessionId: string, agentName: string | undefined, defaultMode?: EditPermissionMode) => { return useContextStore.getState().getSessionAgentEditMode(sessionId, agentName, defaultMode); @@ -340,6 +341,63 @@ export const useSessionStore = create()( return useContextStore.getState().pollForTokenUpdates(sessionId, messageId, messages, maxAttempts); }, updateSession: (session: Session) => useSessionManagementStore.getState().updateSession(session), + + revertToMessage: async (sessionId: string, messageId: string) => { + // Get the message text before reverting + const messages = useMessageStore.getState().messages.get(sessionId) || []; + const targetMessage = messages.find((m) => m.info.id === messageId); + let messageText = ''; + + if (targetMessage && targetMessage.info.role === 'user') { + // Extract text from user message parts + const textParts = targetMessage.parts.filter((p) => p.type === 'text'); + messageText = textParts + .map((p) => { + const part = p as { text?: string; content?: string }; + return part.text || part.content || ''; + }) + .join('\n') + .trim(); + } + + // Call revert API + const updatedSession = await opencodeClient.revertSession(sessionId, messageId); + + // Update session in store (this stores the revert.messageID) + useSessionManagementStore.getState().updateSession(updatedSession); + + // Filter out reverted messages from the store + // Messages with ID >= revert.messageID should be removed + const currentMessages = useMessageStore.getState().messages.get(sessionId) || []; + const revertMessageId = updatedSession.revert?.messageID; + + if (revertMessageId) { + // Find the index of the revert message + const revertIndex = currentMessages.findIndex((m) => m.info.id === revertMessageId); + if (revertIndex !== -1) { + // Keep only messages before the revert point + const filteredMessages = currentMessages.slice(0, revertIndex); + useMessageStore.getState().syncMessages(sessionId, filteredMessages); + } + } + + // Set pending input text for ChatInput to consume + if (messageText) { + set({ pendingInputText: messageText }); + } + }, + + setPendingInputText: (text: string | null) => { + set({ pendingInputText: text }); + }, + + consumePendingInputText: () => { + const text = get().pendingInputText; + if (text !== null) { + set({ pendingInputText: null }); + } + return text; + }, }), { name: "composed-session-store",