diff --git a/TODO.md b/TODO.md index 8d52538..6d44a4d 100644 --- a/TODO.md +++ b/TODO.md @@ -69,10 +69,10 @@ does not fan out. Every comparable CLI has this: opencode `/undo` and `/redo`, Claude Code `/rewind` with checkpoints. There is `/resume` here, which restores a session, and nothing that walks one back. -- [ ] Snapshot files before each prompt, capped at the 100 most recent -- [ ] `/undo` restores files, conversation, or both; `/redo` reverses it -- [ ] Say plainly what is not covered: a `bash` command's effects cannot be snapshotted -- [ ] Test: an edit is reverted, and the model's own record of it goes with it +- [x] Snapshot files before each prompt, capped at the 100 most recent (`src/snapshot.ts` hook + `session.ts` per-turn capture, cap 100 via `SnapshotStack`) +- [x] `/undo` restores files, conversation, or both; `/redo` reverses it (`src/commands.ts` + `src/session.ts` `undo()`/`redo()` + `src/ui/App.tsx` — files+messages together, redo replays tail) +- [x] Say plainly what is not covered: a `bash` command's effects cannot be snapshotted (notice in undo/redo output + `src/snapshot.ts` doc) +- [x] Test: an edit is reverted, and the model's own record of it goes with it (`test/undo.test.ts`: undo file+messages, redo file+messages, bash-not-snapshotted, cap 100) --- diff --git a/src/commands.ts b/src/commands.ts index 8004523..efadee5 100644 --- a/src/commands.ts +++ b/src/commands.ts @@ -26,6 +26,8 @@ export type CommandAction = | { type: 'info'; text: string } | { type: 'model'; model: string } | { type: 'resume'; id: string } + | { type: 'undo' } + | { type: 'redo' } /** A custom command from a markdown file, expanded against its arguments. */ | { type: 'custom'; command: CustomCommand; args: string[] } | { type: 'unknown'; name: string }; @@ -61,6 +63,8 @@ export const COMMANDS: CommandSpec[] = [ { name: 'sessions', summary: 'list saved sessions' }, { name: 'resume', arg: '', summary: 'load a saved session' }, { name: 'save', summary: 'write the session to disk now' }, + { name: 'undo', summary: 'undo the last turn — restores files and conversation (bash effects are not snapshotted)' }, + { name: 'redo', summary: 'redo the last undone turn' }, { name: 'clear', summary: 'clear the transcript and history' }, { name: 'exit', aliases: ['quit'], summary: 'quit' }, ]; @@ -225,6 +229,10 @@ export function parseCommand(raw: string, custom: readonly CustomCommand[] = []) return arg ? { type: 'model', model: arg } : { type: 'models' }; case 'resume': return arg ? { type: 'resume', id: arg } : { type: 'info', text: 'usage: /resume ' }; + case 'undo': + return { type: 'undo' }; + case 'redo': + return { type: 'redo' }; default: { const cmd = custom.find((c) => c.name === name); return cmd ? { type: 'custom', command: cmd, args: arg ? arg.split(/\s+/) : [] } : { type: 'unknown', name }; diff --git a/src/session.ts b/src/session.ts index 5dfdd42..50b75fd 100644 --- a/src/session.ts +++ b/src/session.ts @@ -21,6 +21,7 @@ import { detachProviderItems, droppedSpan, estimateTokens as pruneEstimateTokens import { createSkillTool, renderSkills, type Skill } from './skills'; import { suggestSkillsFromTranscript, writeAutoSkill } from './skill-learner'; import { disabledToolNames, onBashOutput, tools as builtinTools, type ToolSetName } from './tools'; +import { onBeforeWrite, SnapshotStack, type FileState } from './snapshot'; export type ApprovalRequest = { approvalId: string; @@ -204,6 +205,9 @@ export class Session { /** The 80% spend warning is shown once, not on every turn past the line. */ private warnedSpend = false; private controller: AbortController | undefined; + private readonly snapshots = new SnapshotStack(); + private turnBeforeLen = 0; + private turnBeforeFiles = new Map(); private learnTurns = 0; private lastLearnLen = 0; @@ -351,6 +355,7 @@ export class Session { this.subagentOutputTokens = 0; this.warnedSpend = false; this.notebook.clear(); + this.snapshots.clear(); this.opts.onChange?.(this.messages); } @@ -379,6 +384,54 @@ export class Session { return this.opts.compactThreshold ?? DEFAULT_COMPACT_THRESHOLD; } + canUndo(): boolean { return this.snapshots.canUndo(); } + canRedo(): boolean { return this.snapshots.canRedo(); } + + async undo(): Promise { + const snap = this.snapshots.popForUndo(); + if (!snap) throw new Error('nothing to undo'); + await this.restoreFiles(snap.beforeFiles); + // truncate messages to beforeLen; the tail is kept inside snap for redo + this.messages.length = snap.beforeLen; + this.opts.onChange?.(this.messages); + const n = snap.beforeFiles.size; + const filesNote = n === 0 ? 'no files to restore' : `${n} file(s) restored`; + const msgNote = snap.afterLen > snap.beforeLen ? `${snap.afterLen - snap.beforeLen} message(s) removed` : 'no messages to remove'; + return `undone: ${filesNote}, ${msgNote} (bash effects, if any, were not snapshotted)`; + } + + async redo(): Promise { + const snap = this.snapshots.popForRedo(); + if (!snap) throw new Error('nothing to redo'); + await this.restoreFiles(snap.afterFiles); + // redo replays the tail that undo removed; stored in afterFiles? we also need messages tail. + // The messages tail is the slice that was removed on undo; reconstruct by re-inserting from stored span is not enough + // because snapshots hold beforeLen/afterLen but not the actual messages content. + // We store the removed tail inside the snapshot at push time as an extra field via (snap as any)._tail. + const tail = (snap as unknown as { _tail?: import('ai').ModelMessage[] })._tail; + if (tail && tail.length > 0) { + this.messages.push(...tail); + this.opts.onChange?.(this.messages); + } + const n = snap.afterFiles.size; + const filesNote = n === 0 ? 'no files to restore' : `${n} file(s) restored`; + return `redone: ${filesNote} (bash effects, if any, were not snapshotted)`; + } + + private async restoreFiles(state: Map): Promise { + for (const [abs, st] of state) { + try { + if (!st.existed) { + if (await Bun.file(abs).exists()) await Bun.file(abs).delete(); + } else { + await Bun.write(abs, st.content ?? ''); + } + } catch { + // best-effort per file; one failure should not stop the rest + } + } + } + /** * The session's spend so far and the configured ceiling, for the UI's status * and the refuse-the-next-turn check. Unpriced models report no spend: a @@ -535,6 +588,18 @@ export class Session { return; } + // snapshot boundary: remember messages length before this turn and arm file capture + this.turnBeforeLen = this.messages.length; + this.turnBeforeFiles = new Map(); + onBeforeWrite(async (abs: string) => { + if (this.turnBeforeFiles.has(abs)) return; + const exists = await Bun.file(abs).exists(); + let content: string | null = null; + if (exists) { + try { content = await Bun.file(abs).text(); } catch { content = null; } + } + this.turnBeforeFiles.set(abs, { existed: exists, content }); + }); this.messages.push({ role: 'user', content: userText }); this.opts.onChange?.(this.messages); this.controller = new AbortController(); @@ -551,9 +616,39 @@ export class Session { this.opts.onToolOutput?.(toolCallId, chunk); }); + let turnFailed = false; try { yield* this.run(signal, threshold, outputs); + } catch (e) { + turnFailed = true; + throw e; } finally { + onBeforeWrite(undefined); + // finalize snapshot only for turns that actually ran (even if they errored after writing files, + // the file state is still worth snapshotting so undo can revert a half-failed turn) + try { + if (this.turnBeforeFiles.size > 0 || this.messages.length > this.turnBeforeLen) { + const afterFiles = new Map(); + for (const abs of this.turnBeforeFiles.keys()) { + const exists = await Bun.file(abs).exists(); + let content: string | null = null; + if (exists) { try { content = await Bun.file(abs).text(); } catch { content = null; } } + afterFiles.set(abs, { existed: exists, content }); + } + // tail for redo: the messages added by this turn + const tail = this.messages.slice(this.turnBeforeLen).map((m) => ({ ...m, content: typeof m.content === 'string' ? m.content : JSON.parse(JSON.stringify(m.content)) } as import('ai').ModelMessage)); + const snap: import('./snapshot').TurnSnapshot & { _tail?: import('ai').ModelMessage[] } = { + beforeLen: this.turnBeforeLen, + afterLen: this.messages.length, + beforeFiles: new Map(this.turnBeforeFiles), + afterFiles, + }; + (snap as unknown as { _tail?: import('ai').ModelMessage[] })._tail = tail; + // even failed turns push so undo can revert the file side; empty no-op turns are skipped above + if (!turnFailed || this.turnBeforeFiles.size > 0) this.snapshots.push(snap); + } + } catch {} + this.turnBeforeFiles = new Map(); this.controller = undefined; this.drainPendingHotReload(); onBashOutput(undefined); diff --git a/src/snapshot.ts b/src/snapshot.ts new file mode 100644 index 0000000..623bad5 --- /dev/null +++ b/src/snapshot.ts @@ -0,0 +1,72 @@ +/** + * Per-turn file snapshots for /undo and /redo. + * + * Bash is intentionally not snapshotted: a shell command can do anything + * (network, database, chmod, rm -rf) and there is no way to know what to + * restore. The docs and the undo notice say so plainly. + * + * File tools call `recordBeforeWrite(abs)` before their first write to a path + * in the current turn. Session drains the map at turn boundaries into its + * history stack (cap 100) and owns undo/redo. + */ + +export type FileState = { existed: boolean; content: string | null }; + +export type TurnSnapshot = { + /** Messages length before the turn's user message was pushed. */ + beforeLen: number; + /** Messages length after the turn completed (including tool results). */ + afterLen: number; + /** File state before the turn, keyed by absolute path. Only files the turn touched. */ + beforeFiles: Map; + /** File state after the turn, for redo. */ + afterFiles: Map; +}; + +const MAX_HISTORY = 100; + +let hook: ((abs: string) => Promise | void) | undefined; + +export function onBeforeWrite(fn: ((abs: string) => Promise | void) | undefined): void { + hook = fn; +} + +export async function recordBeforeWrite(abs: string): Promise { + const fn = hook; + if (fn) await fn(abs); +} + +export class SnapshotStack { + private readonly history: TurnSnapshot[] = []; + private readonly future: TurnSnapshot[] = []; + + push(entry: TurnSnapshot): void { + this.history.push(entry); + if (this.history.length > MAX_HISTORY) this.history.shift(); + this.future.length = 0; + } + + canUndo(): boolean { return this.history.length > 0; } + canRedo(): boolean { return this.future.length > 0; } + + popForUndo(): TurnSnapshot | undefined { + const e = this.history.pop(); + if (e) this.future.push(e); + return e; + } + + popForRedo(): TurnSnapshot | undefined { + const e = this.future.pop(); + if (e) this.history.push(e); + return e; + } + + clear(): void { + this.history.length = 0; + this.future.length = 0; + } + + depth(): { undo: number; redo: number } { + return { undo: this.history.length, redo: this.future.length }; + } +} diff --git a/src/tools-extra.ts b/src/tools-extra.ts index ab8a77d..3b58a8e 100644 --- a/src/tools-extra.ts +++ b/src/tools-extra.ts @@ -3,6 +3,7 @@ import { stat } from 'node:fs/promises'; import { resolve } from 'node:path'; import { z } from 'zod'; import { jail, posix, walk } from './ignore'; +import { recordBeforeWrite } from './snapshot'; import { withMeta } from './tool-utils'; import { git } from './tools-git'; @@ -47,6 +48,7 @@ export const insertLinesTool = withMeta({ set: 'extra', mutating: true }, tool({ }), execute: async ({ path, line, text }) => { const { abs, lines: cur } = await readLines(path); + await recordBeforeWrite(abs); if (line > cur.length + 1) throw new Error(`line ${line} is past the end of ${path} (${cur.length} lines)`); cur.splice(line - 1, 0, ...lines(text)); await Bun.write(abs, cur.join('\n')); @@ -64,6 +66,7 @@ export const deleteLinesTool = withMeta({ set: 'extra', mutating: true }, tool({ execute: async ({ path, start, end }) => { if (end < start) throw new Error('end must be >= start'); const { abs, lines: cur } = await readLines(path); + await recordBeforeWrite(abs); if (end > cur.length) throw new Error(`end ${end} is past the end of ${path} (${cur.length} lines)`); if (start === 1 && end === cur.length) throw new Error('that deletes the whole file; use delete_file instead'); cur.splice(start - 1, end - start + 1); @@ -83,6 +86,7 @@ export const replaceLinesTool = withMeta({ set: 'extra', mutating: true }, tool( execute: async ({ path, start, end, text }) => { if (end < start) throw new Error('end must be >= start'); const { abs, lines: cur } = await readLines(path); + await recordBeforeWrite(abs); if (end > cur.length) throw new Error(`end ${end} is past the end of ${path} (${cur.length} lines)`); cur.splice(start - 1, end - start + 1, ...lines(text)); await Bun.write(abs, cur.join('\n')); @@ -95,6 +99,7 @@ export const appendFileTool = withMeta({ set: 'extra', mutating: true }, tool({ inputSchema: z.object({ path: z.string(), text: z.string() }), execute: async ({ path, text }) => { const { abs, lines: cur } = await readLines(path); + await recordBeforeWrite(abs); await Bun.write(abs, `${cur.join('\n').replace(/\n?$/, '\n')}${text.replace(/\n?$/, '')}\n`); return `Appended ${lines(text).length} line(s) to ${path}`; }, @@ -105,6 +110,7 @@ export const prependFileTool = withMeta({ set: 'extra', mutating: true }, tool({ inputSchema: z.object({ path: z.string(), text: z.string() }), execute: async ({ path, text }) => { const { abs, lines: cur } = await readLines(path); + await recordBeforeWrite(abs); await Bun.write(abs, `${text.replace(/\n?$/, '\n')}${cur.join('\n')}`); return `Prepended ${lines(text).length} line(s) to ${path}`; }, diff --git a/src/tools.ts b/src/tools.ts index 96dc4b9..8b430a0 100644 --- a/src/tools.ts +++ b/src/tools.ts @@ -3,6 +3,7 @@ import { stat } from 'node:fs/promises'; import { join, resolve } from 'node:path'; import { z } from 'zod'; import { jail, posix, walk } from './ignore'; +import { recordBeforeWrite } from './snapshot'; import { EXTRA_TOOL_NAMES, extraTools } from './tools-extra'; import { GIT_TOOL_NAMES, gitTools } from './tools-git'; import { NET_TOOL_NAMES, netTools } from './tools-net'; @@ -206,6 +207,8 @@ export const applyPatchTool = withMeta({ set: 'edit-plus', mutating: true }, too if (seen.has(op.path)) throw new Error(`${op.path} appears twice in one patch`); seen.add(op.path); const abs = jail(op.path); + await recordBeforeWrite(abs); + if ((op as { moveTo?: string }).moveTo) await recordBeforeWrite(jail((op as { moveTo?: string }).moveTo!)); if (op.kind === 'delete') { if (!(await Bun.file(abs).exists())) throw new Error(`cannot delete ${op.path}: no such file`); @@ -274,6 +277,7 @@ export const writeFileTool = withMeta({ set: 'core', mutating: true }, tool({ }), execute: async ({ path, content }) => { const abs = jail(path); + await recordBeforeWrite(abs); const before = await Bun.file(abs).exists() ? await Bun.file(abs).text() : undefined; await Bun.write(abs, content); @@ -300,6 +304,7 @@ export const editFileTool = withMeta({ set: 'core', mutating: true }, tool({ execute: async ({ path, oldString, newString, replaceAll = false }) => { if (oldString === newString) throw new Error('oldString and newString are identical'); const abs = jail(path); + await recordBeforeWrite(abs); const file = Bun.file(abs); if (!(await file.exists())) throw new Error(`No such file: ${path}`); const before = await file.text(); @@ -337,6 +342,7 @@ export const multiEditTool = withMeta({ set: 'edit-plus', mutating: true }, tool }), execute: async ({ path, edits }) => { const abs = jail(path); + await recordBeforeWrite(abs); const file = Bun.file(abs); if (!(await file.exists())) throw new Error(`No such file: ${path}`); @@ -700,6 +706,8 @@ export const moveFileTool = withMeta({ set: 'edit-plus', mutating: true }, tool( execute: async ({ from, to }) => { const source = jail(from); const target = jail(to); + await recordBeforeWrite(source); + await recordBeforeWrite(target); if (source === target) throw new Error('from and to are the same path'); const file = Bun.file(source); @@ -721,6 +729,7 @@ export const deleteFileTool = withMeta({ set: 'edit-plus', mutating: true }, too }), execute: async ({ path }) => { const abs = jail(path); + await recordBeforeWrite(abs); // Bun.file on a directory reports exists() false, so the stat is what // distinguishes "missing" from "a directory" and gives the right refusal. diff --git a/src/ui/App.tsx b/src/ui/App.tsx index 523af02..6dd30a7 100644 --- a/src/ui/App.tsx +++ b/src/ui/App.tsx @@ -674,6 +674,28 @@ export function App({ push({ kind: 'error', text: e instanceof Error ? e.message : String(e) }); } return; + case 'undo': { + push({ kind: 'user', text: chosen.trim() }); + setWorking(true); + try { + push({ kind: 'info', text: await session.undo() }); + } catch (e) { + push({ kind: 'error', text: e instanceof Error ? e.message : String(e) }); + } + setWorking(false); + return; + } + case 'redo': { + push({ kind: 'user', text: chosen.trim() }); + setWorking(true); + try { + push({ kind: 'info', text: await session.redo() }); + } catch (e) { + push({ kind: 'error', text: e instanceof Error ? e.message : String(e) }); + } + setWorking(false); + return; + } case 'provider': push({ kind: 'user', text: chosen.trim() }); setOnboarding(true); diff --git a/test/undo.test.ts b/test/undo.test.ts new file mode 100644 index 0000000..9ed59ba --- /dev/null +++ b/test/undo.test.ts @@ -0,0 +1,112 @@ +import { expect, test } from 'bun:test'; +import { MockLanguageModelV4, simulateReadableStream } from 'ai/test'; +import type { LanguageModelV4StreamPart } from '@ai-sdk/provider'; +import { mkdtempSync, rmSync } from 'node:fs'; +import { tmpdir } from 'node:os'; +import { join } from 'node:path'; +import { Session } from '../src/session'; + +const usage = { inputTokens: { total: 10, noCache: 10, cacheRead: 0, cacheWrite: 0 }, outputTokens: { total: 5, text: 5, reasoning: 0 } } as unknown as import('@ai-sdk/provider').LanguageModelV4Usage; + +function stream(parts: LanguageModelV4StreamPart[]) { + return { stream: simulateReadableStream({ chunks: parts, chunkDelayInMs: null, initialDelayInMs: null }) }; +} +function toolCall(id: string, toolName: string, input: unknown): LanguageModelV4StreamPart[] { + return [ + { type: 'tool-input-start', id, toolName }, + { type: 'tool-input-end', id }, + { type: 'tool-call', toolCallId: id, toolName, input: JSON.stringify(input) }, + { type: 'finish', finishReason: { unified: 'tool-calls', raw: 'tool_use' }, usage }, + ]; +} +function text(body: string): LanguageModelV4StreamPart[] { + return [ + { type: 'text-start', id: '0' }, + { type: 'text-delta', id: '0', delta: body }, + { type: 'text-end', id: '0' }, + { type: 'finish', finishReason: { unified: 'stop', raw: 'stop' }, usage }, + ]; +} +function inTempDir(fn: () => Promise): Promise { + const orig = process.cwd(); + const dir = mkdtempSync(join(tmpdir(), 'shiro-undo-')); + process.chdir(dir); + return fn().finally(() => { + process.chdir(orig); + rmSync(dir, { recursive: true, force: true }); + }); +} + +test('undo restores file and removes the turn messages', async () => + inTempDir(async () => { + const p = join(process.cwd(), 'note.txt'); + await Bun.write(p, 'before\n'); + let call = 0; + const session = new Session({ + yolo: true, + model: new MockLanguageModelV4({ + doStream: async () => stream(call++ === 0 ? toolCall('c1', 'write_file', { path: 'note.txt', content: 'after\n' }) : text('done')), + }), + askApproval: async () => 'once', + }); + for await (const _ of session.send('overwrite note')) void _; + expect(await Bun.file(p).text()).toBe('after\n'); + const lenAfter = session.messages.length; + const msg = await session.undo(); + expect(msg).toMatch(/undone/); + expect(await Bun.file(p).text()).toBe('before\n'); + expect(session.messages.length).toBeLessThan(lenAfter); + })); + +test('redo restores file and messages after undo', async () => + inTempDir(async () => { + const p = join(process.cwd(), 'note.txt'); + await Bun.write(p, 'before\n'); + let call = 0; + const session = new Session({ + yolo: true, + model: new MockLanguageModelV4({ + doStream: async () => stream(call++ === 0 ? toolCall('c1', 'write_file', { path: 'note.txt', content: 'after\n' }) : text('done')), + }), + askApproval: async () => 'once', + }); + for await (const _ of session.send('overwrite')) void _; + const lenAfter = session.messages.length; + await session.undo(); + expect(await Bun.file(p).text()).toBe('before\n'); + const msg = await session.redo(); + expect(msg).toMatch(/redone/); + expect(await Bun.file(p).text()).toBe('after\n'); + expect(session.messages.length).toBe(lenAfter); + })); + +test('bash effects are not snapshotted (file left, messages still undone)', async () => + inTempDir(async () => { + const p = join(process.cwd(), 'out.txt'); + let call = 0; + const session = new Session({ + yolo: true, + model: new MockLanguageModelV4({ + doStream: async () => stream(call++ === 0 ? toolCall('c1', 'bash', { command: 'echo hi > out.txt' }) : text('done')), + }), + askApproval: async () => 'once', + }); + for await (const _ of session.send('make file via bash')) void _; + expect(await Bun.file(p).exists()).toBe(true); + const lenAfter = session.messages.length; + const msg = await session.undo(); + expect(msg).toMatch(/bash effects.*not snapshotted/); + // bash file remains (not part of snapshot) + expect(await Bun.file(p).exists()).toBe(true); + // but messages are still rewound + expect(session.messages.length).toBeLessThan(lenAfter); + })); + +test('cap 100: oldest snapshot drops', async () => { + const { SnapshotStack } = await import('../src/snapshot'); + const s = new SnapshotStack(); + for (let i = 0; i < 105; i++) { + s.push({ beforeLen: i, afterLen: i + 1, beforeFiles: new Map(), afterFiles: new Map() }); + } + expect(s.depth().undo).toBe(100); +});