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
This commit is contained in:
Bohdan Triapitsyn
2025-12-21 00:47:43 +02:00
parent 70c9a61189
commit ac53ec0bdc
8 changed files with 227 additions and 48 deletions
@@ -71,6 +71,8 @@ export const ChatInput: React.FC<ChatInputProps> = ({ 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<ChatInputProps> = ({ 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;
@@ -66,6 +66,7 @@ interface ChatMessageProps {
scrollToBottom?: (options?: { instant?: boolean; force?: boolean }) => void;
isPendingAnchor?: boolean;
turnGroupingContext?: TurnGroupingContext;
isFirstMessage?: boolean;
}
const ChatMessage: React.FC<ChatMessageProps> = ({
@@ -76,6 +77,7 @@ const ChatMessage: React.FC<ChatMessageProps> = ({
animationHandlers,
isPendingAnchor = false,
turnGroupingContext,
isFirstMessage = false,
}) => {
const { isMobile, hasTouchInput } = useDeviceInfo();
const { currentTheme } = useThemeSystem();
@@ -529,6 +531,13 @@ const ChatMessage: React.FC<ChatMessageProps> = ({
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<ChatMessageProps> = ({
showReasoningTraces={showReasoningTraces}
onAuxiliaryContentComplete={handleAuxiliaryContentComplete}
agentMention={agentMention}
onRevert={handleRevert}
isFirstMessage={isFirstMessage}
/>
</div>
</div>
+20 -13
View File
@@ -75,19 +75,26 @@ const MessageList: React.FC<MessageListProps> = ({
)}
<div className="flex flex-col">
{displayMessages.map((message, index) => (
<ChatMessage
key={message.info.id}
message={message}
previousMessage={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 (
<ChatMessage
key={message.info.id}
message={message}
previousMessage={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)}
isFirstMessage={isFirstUserMessage}
/>
);
})}
</div>
@@ -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<number | null>(null);
@@ -231,40 +235,63 @@ const UserMessageBody: React.FC<{
})}
</div>
<MessageFilesDisplay files={parts} onShowPopup={onShowPopup} />
{canCopyMessage && hasCopyableText && (
{(canCopyMessage && hasCopyableText) || (onRevert && !isFirstMessage) ? (
<div className={cn(
"mt-1 flex items-center justify-end gap-2 opacity-0 pointer-events-none transition-opacity duration-150 group-hover/message:opacity-100 group-hover/message:pointer-events-auto focus-within:opacity-100 focus-within:pointer-events-auto",
copyHintVisible && "opacity-100 pointer-events-auto"
)}>
<Tooltip delayDuration={1000}>
<TooltipTrigger asChild>
<Button
type="button"
variant="ghost"
size="icon"
data-visible={copyHintVisible || isMessageCopied ? 'true' : undefined}
className="h-8 w-8 text-muted-foreground bg-transparent hover:text-foreground hover:!bg-transparent active:!bg-transparent focus-visible:!bg-transparent focus-visible:ring-2 focus-visible:ring-primary/50"
aria-label="Copy message text"
onPointerDown={(event) => event.stopPropagation()}
onClick={handleCopyButtonClick}
onFocus={() => setCopyHintVisible(true)}
onBlur={() => {
if (!isMessageCopied) {
setCopyHintVisible(false);
}
}}
>
{isMessageCopied ? (
<RiCheckLine className="h-3.5 w-3.5 text-[color:var(--status-success)]" />
) : (
<RiFileCopyLine className="h-3.5 w-3.5" />
)}
</Button>
</TooltipTrigger>
<TooltipContent sideOffset={6}>Copy message</TooltipContent>
</Tooltip>
{onRevert && !isFirstMessage && (
<Tooltip delayDuration={1000}>
<TooltipTrigger asChild>
<Button
type="button"
variant="ghost"
size="icon"
className="h-8 w-8 text-muted-foreground bg-transparent hover:text-foreground hover:!bg-transparent active:!bg-transparent focus-visible:!bg-transparent focus-visible:ring-2 focus-visible:ring-primary/50"
aria-label="Revert to this message"
onPointerDown={(event) => event.stopPropagation()}
onClick={(event) => {
event.stopPropagation();
onRevert();
}}
>
<RiArrowGoBackLine className="h-3.5 w-3.5" />
</Button>
</TooltipTrigger>
<TooltipContent sideOffset={6}>Revert from here</TooltipContent>
</Tooltip>
)}
{canCopyMessage && hasCopyableText && (
<Tooltip delayDuration={1000}>
<TooltipTrigger asChild>
<Button
type="button"
variant="ghost"
size="icon"
data-visible={copyHintVisible || isMessageCopied ? 'true' : undefined}
className="h-8 w-8 text-muted-foreground bg-transparent hover:text-foreground hover:!bg-transparent active:!bg-transparent focus-visible:!bg-transparent focus-visible:ring-2 focus-visible:ring-primary/50"
aria-label="Copy message text"
onPointerDown={(event) => event.stopPropagation()}
onClick={handleCopyButtonClick}
onFocus={() => setCopyHintVisible(true)}
onBlur={() => {
if (!isMessageCopied) {
setCopyHintVisible(false);
}
}}
>
{isMessageCopied ? (
<RiCheckLine className="h-3.5 w-3.5 text-[color:var(--status-success)]" />
) : (
<RiFileCopyLine className="h-3.5 w-3.5" />
)}
</Button>
</TooltipTrigger>
<TooltipContent sideOffset={6}>Copy message</TooltipContent>
</Tooltip>
)}
</div>
)}
) : null}
</div>
);
};
@@ -1193,6 +1220,8 @@ const MessageBody: React.FC<MessageBodyProps> = ({ isUser, ...props }) => {
copiedMessage={props.copiedMessage}
onShowPopup={props.onShowPopup}
agentMention={props.agentMention}
onRevert={props.onRevert}
isFirstMessage={props.isFirstMessage}
/>
);
}
+19
View File
@@ -478,6 +478,25 @@ class OpencodeService {
return Boolean(response.data);
}
async revertSession(sessionId: string, messageId: string, partId?: string): Promise<Session> {
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<Session> {
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<string, { type: "idle" | "busy" | "retry"; attempt?: number; message?: string; next?: number }>
> {
+37 -4
View File
@@ -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 <T>(sessionId: string | null | undefined, operation: () => Promise<T>): Promise<T> => {
const directoryOverride = await resolveSessionDirectory(sessionId);
if (directoryOverride) {
@@ -358,15 +382,20 @@ export const useMessageStore = create<MessageStore>()(
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<MessageStore>()(
},
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);
@@ -96,6 +96,8 @@ export interface SessionStore {
userSummaryTitles: Map<string, { title: string; createdAt: number | null }>;
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<void>;
setPendingInputText: (text: string | null) => void;
consumePendingInputText: () => string | null;
}
+58
View File
@@ -94,6 +94,7 @@ export const useSessionStore = create<SessionStore>()(
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<SessionStore>()(
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",