diff --git a/apps/server/src/orchestration/Layers/CheckpointReactor.test.ts b/apps/server/src/orchestration/Layers/CheckpointReactor.test.ts index 01430a201289..07fb9cdcc0cd 100644 --- a/apps/server/src/orchestration/Layers/CheckpointReactor.test.ts +++ b/apps/server/src/orchestration/Layers/CheckpointReactor.test.ts @@ -1683,86 +1683,132 @@ describe("CheckpointReactor", () => { }), ); - it("executes provider revert and emits thread.reverted for checkpoint revert requests", async () => { - const harness = await createHarness(); - const createdAt = "2026-01-01T00:00:00.000Z"; + it.each([ + { commandType: "thread.checkpoint.revert", initializeGit: true }, + { commandType: "thread.conversation.revert", initializeGit: true }, + { commandType: "thread.conversation.revert", initializeGit: false }, + ] as const)( + "$commandType rewinds history with the requested filesystem behavior (git: $initializeGit)", + async ({ commandType, initializeGit }) => { + const harness = await createHarness({ + initializeGit, + seedFilesystemCheckpoints: initializeGit, + }); + const createdAt = "2026-01-01T00:00:00.000Z"; - await Effect.runPromise( - harness.engine.dispatch({ - type: "thread.session.set", - commandId: CommandId.make("cmd-session-set"), - threadId: ThreadId.make("thread-1"), - session: { + await Effect.runPromise( + harness.engine.dispatch({ + type: "thread.session.set", + commandId: CommandId.make("cmd-session-set"), threadId: ThreadId.make("thread-1"), - status: "ready", - providerName: "codex", - runtimeMode: "approval-required", - activeTurnId: null, - lastError: null, - updatedAt: createdAt, - }, - createdAt, - }), - ); + session: { + threadId: ThreadId.make("thread-1"), + status: "ready", + providerName: "codex", + runtimeMode: "approval-required", + activeTurnId: null, + lastError: null, + updatedAt: createdAt, + }, + createdAt, + }), + ); - await Effect.runPromise( - harness.engine.dispatch({ - type: "thread.turn.diff.complete", - commandId: CommandId.make("cmd-diff-1"), - threadId: ThreadId.make("thread-1"), - turnId: asTurnId("turn-1"), - completedAt: createdAt, - checkpointRef: checkpointRefForThreadTurn(ThreadId.make("thread-1"), 1), - status: "ready", - files: [], - checkpointTurnCount: 1, - createdAt, - }), - ); - await Effect.runPromise( - harness.engine.dispatch({ - type: "thread.turn.diff.complete", - commandId: CommandId.make("cmd-diff-2"), - threadId: ThreadId.make("thread-1"), - turnId: asTurnId("turn-2"), - completedAt: createdAt, - checkpointRef: checkpointRefForThreadTurn(ThreadId.make("thread-1"), 2), - status: "ready", - files: [], - checkpointTurnCount: 2, - createdAt, - }), - ); + await Effect.runPromise( + harness.engine.dispatch({ + type: "thread.turn.diff.complete", + commandId: CommandId.make("cmd-diff-1"), + threadId: ThreadId.make("thread-1"), + turnId: asTurnId("turn-1"), + completedAt: createdAt, + checkpointRef: initializeGit + ? checkpointRefForThreadTurn(ThreadId.make("thread-1"), 1) + : CheckpointRef.make("provider-diff:thread-1:turn-1"), + status: initializeGit ? "ready" : "missing", + files: [], + checkpointTurnCount: 1, + createdAt, + }), + ); + await Effect.runPromise( + harness.engine.dispatch({ + type: "thread.turn.diff.complete", + commandId: CommandId.make("cmd-diff-2"), + threadId: ThreadId.make("thread-1"), + turnId: asTurnId("turn-2"), + completedAt: createdAt, + checkpointRef: initializeGit + ? checkpointRefForThreadTurn(ThreadId.make("thread-1"), 2) + : CheckpointRef.make("provider-diff:thread-1:turn-2"), + status: initializeGit ? "ready" : "missing", + files: [], + checkpointTurnCount: 2, + createdAt, + }), + ); - await Effect.runPromise( - harness.engine.dispatch({ - type: "thread.checkpoint.revert", - commandId: CommandId.make("cmd-revert-request"), - threadId: ThreadId.make("thread-1"), - turnCount: 1, - createdAt, - }), - ); + NodeFS.writeFileSync(NodePath.join(harness.cwd, "README.md"), "staged edit\n"); + if (initializeGit) { + NodeChildProcess.execFileSync("git", ["add", "README.md"], { cwd: harness.cwd }); + } + NodeFS.writeFileSync(NodePath.join(harness.cwd, "README.md"), "unstaged edit\n"); + NodeFS.writeFileSync(NodePath.join(harness.cwd, "scratch.txt"), "untracked edit\n"); + const indexBefore = initializeGit + ? NodeChildProcess.execFileSync("git", ["ls-files", "--stage"], { + cwd: harness.cwd, + encoding: "utf8", + }) + : undefined; - await waitForEvent(harness.engine, (event) => event.type === "thread.reverted"); - const thread = await waitForThread( - harness.readModel, - (entry) => entry.checkpoints.length === 1, - ); + await Effect.runPromise( + harness.engine.dispatch({ + type: commandType, + commandId: CommandId.make("cmd-revert-request"), + threadId: ThreadId.make("thread-1"), + turnCount: 1, + createdAt, + }), + ); - expect(thread.latestTurn?.turnId).toBe("turn-1"); - expect(thread.checkpoints).toHaveLength(1); - expect(thread.checkpoints[0]?.checkpointTurnCount).toBe(1); - expect(harness.provider.rollbackConversation).toHaveBeenCalledTimes(1); - expect(harness.provider.rollbackConversation).toHaveBeenCalledWith({ - threadId: ThreadId.make("thread-1"), - numTurns: 1, - }); - expect(NodeFS.readFileSync(NodePath.join(harness.cwd, "README.md"), "utf8")).toBe("v2\n"); - expect( - gitRefExists(harness.cwd, checkpointRefForThreadTurn(ThreadId.make("thread-1"), 2)), - ).toBe(false); - }); + await waitForEvent(harness.engine, (event) => event.type === "thread.reverted"); + const thread = await waitForThread( + harness.readModel, + (entry) => entry.checkpoints.length === 1, + ); + + expect(thread.latestTurn?.turnId).toBe("turn-1"); + expect(thread.checkpoints).toHaveLength(1); + expect(thread.checkpoints[0]?.checkpointTurnCount).toBe(1); + expect(harness.provider.rollbackConversation).toHaveBeenCalledTimes(1); + expect(harness.provider.rollbackConversation).toHaveBeenCalledWith({ + threadId: ThreadId.make("thread-1"), + numTurns: 1, + }); + expect(NodeFS.readFileSync(NodePath.join(harness.cwd, "README.md"), "utf8")).toBe( + commandType === "thread.conversation.revert" ? "unstaged edit\n" : "v2\n", + ); + if (commandType === "thread.conversation.revert") { + expect(NodeFS.readFileSync(NodePath.join(harness.cwd, "scratch.txt"), "utf8")).toBe( + "untracked edit\n", + ); + if (initializeGit) { + expect( + NodeChildProcess.execFileSync("git", ["ls-files", "--stage"], { + cwd: harness.cwd, + encoding: "utf8", + }), + ).toBe(indexBefore); + } + } + if (initializeGit) { + expect( + gitRefExists(harness.cwd, checkpointRefForThreadTurn(ThreadId.make("thread-1"), 2)), + ).toBe(false); + } else { + expect(NodeFS.existsSync(NodePath.join(harness.cwd, ".git"))).toBe(false); + } + }, + ); it("executes provider revert and emits thread.reverted for claude sessions", async () => { const harness = await createHarness({ providerName: ProviderDriverKind.make("claudeAgent") }); diff --git a/apps/server/src/orchestration/Layers/CheckpointReactor.ts b/apps/server/src/orchestration/Layers/CheckpointReactor.ts index fc1a3f740349..9cfcc2e74915 100644 --- a/apps/server/src/orchestration/Layers/CheckpointReactor.ts +++ b/apps/server/src/orchestration/Layers/CheckpointReactor.ts @@ -704,16 +704,11 @@ const make = Effect.gen(function* () { thread, projects: yield* resolveThreadProjects(thread.projectId), preferSessionRuntime: true, - }); - if (!checkpointCwd) { - yield* appendRevertFailureActivity({ - threadId: event.payload.threadId, - turnCount: event.payload.turnCount, - detail: "Checkpoint workspace is unavailable or is not a git repository.", - createdAt: now, - }).pipe(Effect.catch(() => Effect.void)); - return; - } + }).pipe( + Effect.catch((error) => + event.payload.restoreFiles === false ? Effect.succeed(undefined) : Effect.fail(error), + ), + ); const currentTurnCount = thread.checkpoints.reduce( (maxTurnCount, checkpoint) => Math.max(maxTurnCount, checkpoint.checkpointTurnCount), @@ -730,43 +725,55 @@ const make = Effect.gen(function* () { return; } - const targetCheckpointRef = - event.payload.turnCount === 0 - ? checkpointRefForThreadTurn(event.payload.threadId, 0) - : thread.checkpoints.find( - (checkpoint) => checkpoint.checkpointTurnCount === event.payload.turnCount, - )?.checkpointRef; + yield* providerService.assertConversationRollbackSupported(event.payload.threadId); - if (!targetCheckpointRef) { - yield* appendRevertFailureActivity({ - threadId: event.payload.threadId, - turnCount: event.payload.turnCount, - detail: `Checkpoint ref for turn ${event.payload.turnCount} is unavailable in read model.`, - createdAt: now, - }).pipe(Effect.catch(() => Effect.void)); - return; - } + if (event.payload.restoreFiles !== false) { + if (!checkpointCwd) { + yield* appendRevertFailureActivity({ + threadId: event.payload.threadId, + turnCount: event.payload.turnCount, + detail: "Checkpoint workspace is unavailable or is not a git repository.", + createdAt: now, + }).pipe(Effect.catch(() => Effect.void)); + return; + } - yield* providerService.assertConversationRollbackSupported(event.payload.threadId); + const targetCheckpointRef = + event.payload.turnCount === 0 + ? checkpointRefForThreadTurn(event.payload.threadId, 0) + : thread.checkpoints.find( + (checkpoint) => checkpoint.checkpointTurnCount === event.payload.turnCount, + )?.checkpointRef; + + if (!targetCheckpointRef) { + yield* appendRevertFailureActivity({ + threadId: event.payload.threadId, + turnCount: event.payload.turnCount, + detail: `Checkpoint ref for turn ${event.payload.turnCount} is unavailable in read model.`, + createdAt: now, + }).pipe(Effect.catch(() => Effect.void)); + return; + } - const restored = yield* checkpointStore.restoreCheckpoint({ - cwd: checkpointCwd, - checkpointRef: targetCheckpointRef, - fallbackToHead: event.payload.turnCount === 0, - }); - if (!restored) { - yield* appendRevertFailureActivity({ - threadId: event.payload.threadId, - turnCount: event.payload.turnCount, - detail: `Filesystem checkpoint is unavailable for turn ${event.payload.turnCount}.`, - createdAt: now, - }).pipe(Effect.catch(() => Effect.void)); - return; - } + const restored = yield* checkpointStore.restoreCheckpoint({ + cwd: checkpointCwd, + checkpointRef: targetCheckpointRef, + fallbackToHead: event.payload.turnCount === 0, + }); + if (!restored) { + yield* appendRevertFailureActivity({ + threadId: event.payload.threadId, + turnCount: event.payload.turnCount, + detail: `Filesystem checkpoint is unavailable for turn ${event.payload.turnCount}.`, + createdAt: now, + }).pipe(Effect.catch(() => Effect.void)); + return; + } - // Refresh the workspace entry index so the @-mention file picker - // reflects the reverted filesystem state. - yield* workspaceEntries.refresh(checkpointCwd); + // Refresh the workspace entry index so the @-mention file picker + // reflects the reverted filesystem state. + yield* workspaceEntries.refresh(checkpointCwd); + } const rolledBackTurns = Math.max(0, currentTurnCount - event.payload.turnCount); if (rolledBackTurns > 0) { @@ -783,7 +790,7 @@ const make = Effect.gen(function* () { } } - if (staleCheckpointRefs.length > 0) { + if (checkpointCwd && staleCheckpointRefs.length > 0) { yield* checkpointStore.deleteCheckpointRefs({ cwd: checkpointCwd, checkpointRefs: staleCheckpointRefs, diff --git a/apps/server/src/orchestration/decider.ts b/apps/server/src/orchestration/decider.ts index c0787a18d096..96f5c6e2f7f5 100644 --- a/apps/server/src/orchestration/decider.ts +++ b/apps/server/src/orchestration/decider.ts @@ -1640,6 +1640,7 @@ export const decideOrchestrationCommand = Effect.fn("decideOrchestrationCommand" }; } + case "thread.conversation.revert": case "thread.checkpoint.revert": { yield* requireThread({ readModel, @@ -1657,6 +1658,7 @@ export const decideOrchestrationCommand = Effect.fn("decideOrchestrationCommand" payload: { threadId: command.threadId, turnCount: command.turnCount, + ...(command.type === "thread.conversation.revert" ? { restoreFiles: false } : {}), createdAt: command.createdAt, }, }; diff --git a/apps/server/src/provider/Layers/OpenCodeAdapter.test.ts b/apps/server/src/provider/Layers/OpenCodeAdapter.test.ts index ee5767f9d356..a6636fcd69d0 100644 --- a/apps/server/src/provider/Layers/OpenCodeAdapter.test.ts +++ b/apps/server/src/provider/Layers/OpenCodeAdapter.test.ts @@ -97,6 +97,8 @@ const runtimeMock = { promptEchoEvents: [] as Array, closeError: null as Error | null, messages: [] as MessageEntry[], + forkMessagesBySession: new Map(), + forkPreservesBoundary: true, subscribedEvents: [] as Array>, eventSubscribeObserved: null as (() => void) | null, eventStreamError: null as ((cause: unknown) => void) | null, @@ -128,7 +130,7 @@ const runtimeMock = { permissionListImplementation: null as (() => Promise>) | null, questionListImplementation: null as (() => Promise>) | null, sessionUpdateCalls: [] as Array<{ sessionID: string; permission: unknown }>, - forkCalls: [] as Array<{ sessionID: string; directory?: string }>, + forkCalls: [] as Array<{ sessionID: string; directory?: string; messageID?: string }>, }, reset() { this.state.startCalls.length = 0; @@ -157,6 +159,8 @@ const runtimeMock = { this.state.promptEchoEvents.length = 0; this.state.closeError = null; this.state.messages = []; + this.state.forkMessagesBySession.clear(); + this.state.forkPreservesBoundary = true; this.state.subscribedEvents = []; this.state.eventSubscribeObserved = null; this.state.eventStreamError = null; @@ -264,7 +268,8 @@ const OpenCodeRuntimeTestDouble: OpenCodeRuntimeShape = { return { data: { id: sessionID, - ...(runtimeMock.state.revertMessageID + ...(runtimeMock.state.revertMessageID && + !runtimeMock.state.forkMessagesBySession.has(sessionID) ? { revert: { messageID: runtimeMock.state.revertMessageID } } : {}), ...(directory ? { directory } : {}), @@ -276,10 +281,37 @@ const OpenCodeRuntimeTestDouble: OpenCodeRuntimeShape = { runtimeMock.state.sessionUpdateCalls.push({ sessionID, permission }); return { data: { id: sessionID } }; }, - fork: async ({ sessionID, directory }: { sessionID: string; directory?: string }) => { + fork: async ({ + sessionID, + directory, + messageID, + }: { + sessionID: string; + directory?: string; + messageID?: string; + }) => { // Fork clones history into a new session bound to the directory. const forkedId = `${sessionID}_fork`; - runtimeMock.state.forkCalls.push({ sessionID, ...(directory ? { directory } : {}) }); + runtimeMock.state.forkCalls.push({ + sessionID, + ...(directory ? { directory } : {}), + ...(messageID ? { messageID } : {}), + }); + if (messageID) { + const messages = + runtimeMock.state.forkMessagesBySession.get(sessionID) ?? runtimeMock.state.messages; + const boundary = messages.findIndex((entry) => entry.info.id === messageID); + NodeAssert.notEqual(boundary, -1); + runtimeMock.state.forkMessagesBySession.set( + forkedId, + messages + .slice(0, runtimeMock.state.forkPreservesBoundary ? boundary : messages.length) + .map((entry) => ({ + ...entry, + info: { ...entry.info, id: `${entry.info.id}_fork` }, + })), + ); + } if (directory) { runtimeMock.state.sessionDirectoryById.set(forkedId, directory); } @@ -337,7 +369,10 @@ const OpenCodeRuntimeTestDouble: OpenCodeRuntimeShape = { typeof input.sessionID === "string" && typeof input.messageID === "string" ) { - runtimeMock.state.messages.push({ + const messages = + runtimeMock.state.forkMessagesBySession.get(input.sessionID) ?? + runtimeMock.state.messages; + messages.push({ info: { id: input.messageID, role: "user" }, parts: [], }); @@ -355,14 +390,19 @@ const OpenCodeRuntimeTestDouble: OpenCodeRuntimeShape = { runtimeMock.state.summarizeCalls.push(input); return { data: true }; }, - messages: async () => ({ data: runtimeMock.state.messages }), + messages: async ({ sessionID }: { sessionID: string }) => ({ + data: + runtimeMock.state.forkMessagesBySession.get(sessionID) ?? runtimeMock.state.messages, + }), message: async ({ sessionID, messageID }: { sessionID: string; messageID: string }) => { runtimeMock.state.messageCalls.push({ sessionID, messageID }); if (runtimeMock.state.messageFailures > 0) { runtimeMock.state.messageFailures -= 1; throw new Error("message lookup failed", { cause: { status: 500 } }); } - const message = runtimeMock.state.messages.find((entry) => entry.info.id === messageID); + const messages = + runtimeMock.state.forkMessagesBySession.get(sessionID) ?? runtimeMock.state.messages; + const message = messages.find((entry) => entry.info.id === messageID); if (!message) { throw new Error(`Message not found: ${messageID}`, { cause: { status: 404, body: { name: "NotFoundError" } }, @@ -6333,7 +6373,7 @@ it.layer(OpenCodeAdapterTestLayer)("OpenCodeAdapterLive", (it) => { }).pipe(Effect.provide(adapterLayer)); }); - it.effect("reverts the first removed assistant message and returns only retained turns", () => + it.effect("forks before the removed user prompt and resumes only retained history", () => Effect.gen(function* () { const adapter = yield* OpenCodeAdapter; const threadId = asThreadId("thread-rollback-all"); @@ -6356,49 +6396,116 @@ it.layer(OpenCodeAdapterTestLayer)("OpenCodeAdapterLive", (it) => { }, ]; + const originalCursor = (yield* adapter.listSessions()).find( + (session) => session.threadId === threadId, + )?.resumeCursor; + runtimeMock.state.forkPreservesBoundary = false; + const boundaryError = yield* adapter.rollbackThread(threadId, 1).pipe(Effect.flip); + NodeAssert.match(boundaryError.message, /did not preserve the requested rewind boundary/); + NodeAssert.deepEqual( + (yield* adapter.listSessions()).find((session) => session.threadId === threadId) + ?.resumeCursor, + originalCursor, + ); + runtimeMock.state.forkPreservesBoundary = true; + for (const numTurns of [0, 1, 2, 3]) { - runtimeMock.state.revertMessageID = undefined; - runtimeMock.state.revertCalls.length = 0; + yield* adapter.stopSession(threadId); + yield* adapter.startSession({ threadId, runtimeMode: "full-access" }); + runtimeMock.state.forkCalls.length = 0; const snapshot = yield* adapter.rollbackThread(threadId, numTurns); NodeAssert.deepEqual( - runtimeMock.state.revertCalls, + runtimeMock.state.forkCalls.map(({ sessionID, messageID }) => ({ sessionID, messageID })), numTurns === 0 ? [] : [ { sessionID: "http://127.0.0.1:9999/session", - messageID: numTurns === 1 ? "assistant-2" : "assistant-1", + messageID: numTurns === 1 ? "user-2" : "user-1", }, ], ); NodeAssert.deepEqual( snapshot.turns.map((turn) => turn.id), - ["assistant-1", "assistant-2"].slice(0, Math.max(0, 2 - numTurns)), + numTurns === 0 + ? ["assistant-1", "assistant-2"] + : ["assistant-1_fork"].slice(0, Math.max(0, 2 - numTurns)), ); + NodeAssert.deepEqual(runtimeMock.state.revertCalls, []); } - runtimeMock.state.revertMessageID = undefined; + yield* adapter.stopSession(threadId); + yield* adapter.startSession({ threadId, runtimeMode: "full-access" }); for (const remaining of [1, 0]) { const snapshot = yield* adapter.rollbackThread(threadId, 1); NodeAssert.equal(snapshot.turns.length, remaining); NodeAssert.deepEqual((yield* adapter.readThread(threadId)).turns, snapshot.turns); + const cursor = (yield* adapter.listSessions()).find( + (session) => session.threadId === threadId, + )?.resumeCursor; + NodeAssert.deepEqual(cursor, { + schemaVersion: 1, + sessionId: + remaining === 1 + ? "http://127.0.0.1:9999/session_fork" + : "http://127.0.0.1:9999/session_fork_fork", + }); + yield* adapter.stopSession(threadId); + yield* adapter.startSession({ threadId, runtimeMode: "full-access", resumeCursor: cursor }); + NodeAssert.deepEqual((yield* adapter.readThread(threadId)).turns, snapshot.turns); } NodeAssert.deepEqual( - runtimeMock.state.revertCalls.slice(-2).map((call) => call.messageID), - ["assistant-2", "assistant-1"], + runtimeMock.state.forkCalls.slice(-2).map((call) => call.messageID), + ["user-2", "user-1_fork"], ); - runtimeMock.state.revertMessageID = undefined; + yield* adapter.sendTurn({ + threadId, + input: "continue the retained conversation", + modelSelection: createModelSelection( + ProviderInstanceId.make("opencode"), + "anthropic/claude-sonnet-4-5", + ), + }); + NodeAssert.equal( + (runtimeMock.state.promptCalls.at(-1) as { sessionID: string }).sessionID, + "http://127.0.0.1:9999/session_fork_fork", + ); + const continuation = runtimeMock.state.promptCalls.at(-1) as { + sessionID: string; + messageID: string; + }; + runtimeMock.state.forkMessagesBySession.get(continuation.sessionID)!.push({ + info: { id: "continuation-answer", role: "assistant" }, + parts: [{ id: "continuation-part", type: "text", text: "continued answer" }], + }); + const continuationCursor = (yield* adapter.listSessions()).find( + (session) => session.threadId === threadId, + )?.resumeCursor; + yield* adapter.stopSession(threadId); + yield* adapter.startSession({ + threadId, + runtimeMode: "full-access", + resumeCursor: continuationCursor, + }); + NodeAssert.deepEqual( + (yield* adapter.readThread(threadId)).turns.map((turn) => turn.id), + ["continuation-answer"], + ); + NodeAssert.deepEqual((yield* adapter.rollbackThread(threadId, 1)).turns, []); + NodeAssert.equal(runtimeMock.state.forkCalls.at(-1)?.messageID, continuation.messageID); + yield* adapter.stopSession(threadId); + yield* adapter.startSession({ threadId, runtimeMode: "full-access" }); runtimeMock.state.messages = runtimeMock.state.messages.filter( (entry) => entry.info.id !== "user-2", ); const sharedUserSnapshot = yield* adapter.rollbackThread(threadId, 1); - NodeAssert.equal(runtimeMock.state.revertMessageID, "user-1"); + NodeAssert.equal(runtimeMock.state.forkCalls.at(-1)?.messageID, "user-1"); NodeAssert.deepEqual(sharedUserSnapshot.turns, []); NodeAssert.deepEqual((yield* adapter.readThread(threadId)).turns, []); runtimeMock.state.messages = []; - runtimeMock.state.revertCalls.length = 0; + runtimeMock.state.forkCalls.length = 0; const emptySnapshot = yield* adapter.rollbackThread(threadId, 1); - NodeAssert.deepEqual(runtimeMock.state.revertCalls, []); + NodeAssert.deepEqual(runtimeMock.state.forkCalls, []); NodeAssert.deepEqual(emptySnapshot.turns, []); }), ); diff --git a/apps/server/src/provider/Layers/OpenCodeAdapter.ts b/apps/server/src/provider/Layers/OpenCodeAdapter.ts index 312a9494bebc..3d216bb1167b 100644 --- a/apps/server/src/provider/Layers/OpenCodeAdapter.ts +++ b/apps/server/src/provider/Layers/OpenCodeAdapter.ts @@ -338,7 +338,7 @@ interface OpenCodeSessionContext { readonly client: OpencodeClient; readonly server: OpenCodeServerConnection; readonly directory: string; - readonly openCodeSessionId: string; + openCodeSessionId: string; readonly relatedSessionIds: Set; readonly resolvedRequestIds: Set; readonly autoRepliedRequestIds: Set; @@ -3812,14 +3812,89 @@ export function makeOpenCodeAdapter( const targetIndex = Math.max(0, snapshot.turns.length - numTurns); const target = snapshot.turns[targetIndex]; if (target) { - yield* runOpenCodeSdk("session.revert", () => - context.client.session.revert({ + const messages = yield* runOpenCodeSdk("session.messages", () => + context.client.session.messages({ sessionID: context.openCodeSessionId }), + ).pipe(Effect.mapError(toRequestError)); + const entries = messages.data ?? []; + const targetMessageIndex = entries.findIndex((entry) => entry.info.id === target.id); + if (targetMessageIndex < 0) { + return yield* toRequestError( + new OpenCodeRuntimeError({ + operation: "session.fork", + detail: "The OpenCode rewind boundary is no longer available.", + }), + ); + } + const firstRemovedMessage = + entries + .slice(0, targetMessageIndex + 1) + .findLast((entry) => entry.info.role === "user") ?? entries[targetMessageIndex]!; + // Native revert also rewrites workspace files. Fork only the retained + // conversation so T3 alone decides whether filesystem changes survive. + const fork = yield* runOpenCodeSdk("session.fork", () => + context.client.session.fork({ sessionID: context.openCodeSessionId, - messageID: target.id, + messageID: firstRemovedMessage.info.id, + directory: context.directory, + }), + ).pipe(Effect.mapError(toRequestError)); + if (!fork.data) { + return yield* toRequestError( + new OpenCodeRuntimeError({ + operation: "session.fork", + detail: "OpenCode session.fork returned no session payload.", + }), + ); + } + const forkedSessionId = fork.data.id; + const forkMessages = yield* runOpenCodeSdk("session.messages", () => + context.client.session.messages({ sessionID: forkedSessionId }), + ).pipe(Effect.mapError(toRequestError)); + if (forkMessages.data?.length !== entries.indexOf(firstRemovedMessage)) { + return yield* toRequestError( + new OpenCodeRuntimeError({ + operation: "session.fork", + detail: "OpenCode did not preserve the requested rewind boundary.", + }), + ); + } + yield* runOpenCodeSdk("session.update", () => + context.client.session.update({ + sessionID: forkedSessionId, + permission: buildOpenCodePermissionRules(context.session.runtimeMode), }), ).pipe(Effect.mapError(toRequestError)); - // Native revert can move the boundary to the preceding user message. - return yield* readThread(threadId); + yield* clearPendingOpenCodeRequests(context, { type: "session.fork" }); + context.openCodeSessionId = forkedSessionId; + context.relatedSessionIds.clear(); + context.relatedSessionIds.add(forkedSessionId); + context.messageRoleById.clear(); + context.textPartsByMessageId.clear(); + context.turnTokenUsage = undefined; + context.activeTurnId = undefined; + context.interruptedTurnId = undefined; + context.reconcileIdleStatus = false; + context.awaitingBusyAfterInterruption = false; + context.pendingIdleReconciliation = undefined; + context.session = { + ...context.session, + resumeCursor: { schemaVersion: OPENCODE_RESUME_VERSION, sessionId: forkedSessionId }, + updatedAt: yield* nowIso, + }; + yield* emit({ + ...(yield* buildEventBase({ threadId })), + type: "thread.started", + payload: { providerThreadId: forkedSessionId }, + }); + return { + threadId, + turns: forkMessages.data + .filter((entry) => entry.info.role === "assistant") + .map((entry) => ({ + id: TurnId.make(entry.info.id), + items: [entry.info, ...entry.parts], + })), + }; } return snapshot; diff --git a/apps/web/src/components/ChatView.tsx b/apps/web/src/components/ChatView.tsx index 45207934592f..8e9ecf1a5f8f 100644 --- a/apps/web/src/components/ChatView.tsx +++ b/apps/web/src/components/ChatView.tsx @@ -6533,8 +6533,18 @@ export default function ChatView(props: ChatViewProps) { return () => window.removeEventListener("paste", handler, true); }, [activeThreadId, composerRef]); + const [pendingRevert, setPendingRevert] = useState<{ + turnCount: number; + messageId: MessageId; + routeThreadKey: string; + } | null>(null); + + if (pendingRevert && pendingRevert.routeThreadKey !== routeThreadKey) { + setPendingRevert(null); + } + const onRevertToTurnCount = useCallback( - async (turnCount: number, messageId: MessageId) => { + async (turnCount: number, messageId: MessageId, restoreFiles?: boolean) => { const localApi = readLocalApi(); if (!localApi || !activeThread || isRevertingCheckpoint) return; const message = activeThread.messages.find((message) => message.id === messageId); @@ -6558,15 +6568,8 @@ export default function ChatView(props: ChatViewProps) { setThreadError(activeThread.id, "Interrupt the current turn before reverting checkpoints."); return; } - const confirmed = await localApi.dialogs.confirm( - [ - "Edit from here?", - "Rewind files and chat to before this message.", - "Your prompt and attachments return to the composer.", - ].join("\n"), - { variant: "destructive" }, - ); - if (!confirmed) { + if (restoreFiles === undefined) { + setPendingRevert({ turnCount, messageId, routeThreadKey }); return; } @@ -6599,7 +6602,7 @@ export default function ChatView(props: ChatViewProps) { await waitForRevertedMessage(routeThreadRef, messageId, turnCount, async () => { const result = await revertThreadCheckpoint({ environmentId, - input: { threadId: activeThread.id, turnCount }, + input: { threadId: activeThread.id, turnCount, restoreFiles }, }); if (result._tag === "Failure") throw squashAtomCommandFailure(result); }); @@ -9144,6 +9147,44 @@ export default function ChatView(props: ChatViewProps) { ) : null} + { + if (!open) setPendingRevert(null); + }} + > + + + Edit from here? + + Rewind chat to before this message. Your prompt and attachments return to the + composer. + + + + }>Cancel + + + + + {expandedImage && ( { }).pipe(Effect.provide(TEST_CRYPTO_LAYER)), ); + it.effect("uses a distinct command when keeping workspace changes", () => + Effect.gen(function* () { + const dispatched: ClientOrchestrationCommand[] = []; + const supervisor = yield* makeSupervisor(dispatched); + for (const restoreFiles of [undefined, true, false]) { + yield* revertThreadCheckpoint({ + commandId: CommandId.make("rewind-command"), + threadId: ThreadId.make("thread-1"), + turnCount: 0, + ...(restoreFiles !== undefined ? { restoreFiles } : {}), + createdAt: "2026-06-06T00:01:00.000Z", + }).pipe(Effect.provideService(EnvironmentSupervisor.EnvironmentSupervisor, supervisor)); + } + expect(dispatched.map((command) => command.type)).toEqual([ + "thread.checkpoint.revert", + "thread.checkpoint.revert", + "thread.conversation.revert", + ]); + }).pipe(Effect.provide(TEST_CRYPTO_LAYER)), + ); + it.effect("preserves caller metadata for idempotent queued commands", () => Effect.gen(function* () { const dispatched: ClientOrchestrationCommand[] = []; diff --git a/packages/client-runtime/src/operations/commands.ts b/packages/client-runtime/src/operations/commands.ts index 78c16a025d35..8f313c1632ff 100644 --- a/packages/client-runtime/src/operations/commands.ts +++ b/packages/client-runtime/src/operations/commands.ts @@ -53,7 +53,9 @@ export type InterruptThreadTurnInput = CommandInput<"thread.turn.interrupt">; export type RespondToThreadApprovalInput = CommandInput<"thread.approval.respond">; export type RespondToThreadUserInputInput = CommandInput<"thread.user-input.respond">; export type DismissThreadUserInputInput = CommandInput<"thread.user-input.dismiss">; -export type RevertThreadCheckpointInput = CommandInput<"thread.checkpoint.revert">; +export type RevertThreadCheckpointInput = CommandInput<"thread.checkpoint.revert"> & { + readonly restoreFiles?: boolean; +}; export type StopThreadSessionInput = CommandInput<"thread.session.stop">; type DispatchTag = typeof ORCHESTRATION_WS_METHODS.dispatchCommand; @@ -355,9 +357,10 @@ export const dismissThreadUserInput: (input: DismissThreadUserInputInput) => Com export const revertThreadCheckpoint: (input: RevertThreadCheckpointInput) => CommandEffect = Effect.fn("EnvironmentCommands.revertThreadCheckpoint")(function* (input) { const metadata = yield* timestampedCommandMetadata(input); + const { restoreFiles, ...command } = input; return yield* dispatch({ - ...input, - type: "thread.checkpoint.revert", + ...command, + type: restoreFiles === false ? "thread.conversation.revert" : "thread.checkpoint.revert", commandId: metadata.commandId, createdAt: metadata.createdAt, }); diff --git a/packages/contracts/src/orchestration.ts b/packages/contracts/src/orchestration.ts index 56a29b8f65f0..e58ab6b07bfb 100644 --- a/packages/contracts/src/orchestration.ts +++ b/packages/contracts/src/orchestration.ts @@ -1287,6 +1287,13 @@ const ThreadCheckpointRevertCommand = Schema.Struct({ createdAt: IsoDateTime, }); +// A separate command makes older servers reject history-only rewinds rather than +// ignoring an unfamiliar option and restoring files. +const ThreadConversationRevertCommand = Schema.Struct({ + ...ThreadCheckpointRevertCommand.fields, + type: Schema.Literal("thread.conversation.revert"), +}); + const ThreadSessionStopCommand = Schema.Struct({ type: Schema.Literal("thread.session.stop"), commandId: CommandId, @@ -1327,6 +1334,7 @@ const DispatchableClientOrchestrationCommand = Schema.Union([ ThreadUserInputRespondCommand, ThreadUserInputDismissCommand, ThreadCheckpointRevertCommand, + ThreadConversationRevertCommand, ThreadSessionStopCommand, ]); export type DispatchableClientOrchestrationCommand = @@ -1359,6 +1367,7 @@ export const ClientOrchestrationCommand = Schema.Union([ ThreadUserInputRespondCommand, ThreadUserInputDismissCommand, ThreadCheckpointRevertCommand, + ThreadConversationRevertCommand, ThreadSessionStopCommand, ]); export type ClientOrchestrationCommand = typeof ClientOrchestrationCommand.Type; @@ -1762,6 +1771,7 @@ const ThreadUserInputResponseRequestedPayload = Schema.Struct({ export const ThreadCheckpointRevertRequestedPayload = Schema.Struct({ threadId: ThreadId, turnCount: NonNegativeInt, + restoreFiles: Schema.optional(Schema.Boolean), createdAt: IsoDateTime, });